Sentence Transformers 训练总览:模型、数据、损失与现代 Trainer

同一对文本,在不同任务里可能需要不同的相似度。比如“苹果发布新 iPad”和“NVIDIA 准备下一代 GPU”,新闻分类可以把它们都归入科技类,语义相似度模型却应辨认两者表达的事件不同。检索模型主要学习查询与文档之间的匹配关系,也不必把文档之间的距离直接当成目标。针对自己的任务微调嵌入模型,就是让表示空间符合这种需求。

训练流程由四至六个组成部分构成:模型、数据集、损失函数、训练参数、评估器和 Trainer。评估数据与评估器可选,但应在设计训练时就确定如何验证结果。下面依次说明这些组件,以及单数据集、多数据集和旧训练接口的完整例子。

模型:从现有嵌入模型到多模态组合

Sentence Transformer 是模块序列,可以采用内置模块或自定义模块。已训练的 SentenceTransformer 检查点通常带有 modules.json,加载时会恢复原有结构,无须自己猜测池化方式。

from sentence_transformers import SentenceTransformer

model = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2")

如果从普通 Transformer 检查点开始训练,常见结构是 Transformer 加 Pooling:前者得到 token 表示,后者将整句的 token 表示汇总为一个向量。以 BERT 和平均池化为例:

from sentence_transformers import SentenceTransformer
from sentence_transformers.sentence_transformer.modules import Transformer, Pooling

# Loading in fp32 is preferred for training if your memory can handle it
transformer = Transformer("google-bert/bert-base-uncased", model_kwargs={"torch_dtype": "float32"})
pooling = Pooling(transformer.get_embedding_dimension(), pooling_mode="mean")

model = SentenceTransformer(modules=[transformer, pooling])

这种结构是默认路径,也可以使用以下简写。例子保留原文的单精度加载设置,训练时能否采用混合精度,另由训练参数和硬件决定。

from sentence_transformers import SentenceTransformer

model = SentenceTransformer("google-bert/bert-base-uncased", model_kwargs={"torch_dtype": "float32"})

选择基座时,应关注目标语言和领域,而不能只看通用分类榜单。编码器模型通常适合产生 token 或文本向量;可从 Hugging Face 的 fill-mask、sentence-similarity 和 feature-extraction 类别寻找候选。原文以土耳其语为例,建议优先考虑覆盖该语言的 XLM-RoBERTa,而不是只面向英语的 BERT。具体选择仍需用自己的任务数据验证。

静态嵌入

StaticEmbedding 不使用 Transformer 注意力计算,而是按 token 查找已有向量。它减少计算开销,但每个 token 的表示不依赖上下文,因此不能像上下文编码器一样处理复杂语义。以下结构创建一个 512 维静态嵌入模块;它是结构初始化示例,不表示这些新向量已经完成任务训练。

from sentence_transformers import SentenceTransformer
from sentence_transformers.sentence_transformer.modules import StaticEmbedding
from tokenizers import Tokenizer
# Load any Tokenizer from Hugging Face
tokenizer = Tokenizer.from_pretrained("google-bert/bert-base-uncased")
# The `embedding_dim` is the dimensionality (size) of the token embeddings
static_embedding = StaticEmbedding(tokenizer, embedding_dim=512)

model = SentenceTransformer(modules=[static_embedding])

也可以通过 StaticEmbedding.from_model2vec 或 from_distillation 使用现成模型或蒸馏结果。速度、内存和质量应一起评估,本文没有进行性能实测。

视觉语言模型

多模态训练需要相应依赖,例如图像支持可按安装文档使用 pip install -U "sentence-transformers[image]"。Transformer 会读取模型 processor,以识别文本、图像、音频、视频等支持的模态。

可加载已经针对多模态嵌入训练的模型:

from sentence_transformers import SentenceTransformer

model = SentenceTransformer(
    "Qwen/Qwen3-VL-Embedding-2B",
    model_kwargs={"attn_implementation": "flash_attention_2", "torch_dtype": "bfloat16"},
    processor_kwargs={"min_pixels": 28 * 28, "max_pixels": 600 * 600},
)

也可从尚未针对嵌入任务训练的视觉语言模型检查点开始:

from sentence_transformers import SentenceTransformer

model = SentenceTransformer("Qwen/Qwen3-VL-2B")

两条路径都会根据 processor 检查模态,必要时自动添加池化模块。通过属性和方法查看模型能力:

print(model.modalities)
# ['text', 'image', 'video', 'message']

print(model.supports("image"))
# True

这些注释是原文示例输出,并非本次运行所得。message 模态指通过聊天模板组织混合输入,例如同一条输入包含图片和文字。训练数据可以使用文本、PIL 图片、图片路径或 URL、音频数组,以及同时包含多种模态的字典。是否支持某种输入,取决于具体模型和 processor。

上述示例保留了原文的 FlashAttention 2、BF16 和像素范围设置。它们需要匹配的软件与硬件,不是所有 CPU 或 GPU 环境都能直接使用。

使用 Router 组合不同编码器

另一条路线是为每种模态配置独立编码器,再由 Router 分流。下面把 MiniLM 文本编码器和 SigLIP 图像编码器组合起来,并用 Dense 将文本表示投影到图像编码器使用的 768 维空间。SigLIP 自身产生池化表示,因此图像侧没有另加 Pooling。

from sentence_transformers import SentenceTransformer
from sentence_transformers.sentence_transformer.modules import Dense, Pooling, Router, Transformer

# Create separate encoders for different modalities
text_encoder = Transformer("sentence-transformers/all-MiniLM-L6-v2")
text_pooling = Pooling(text_encoder.get_embedding_dimension(), pooling_mode="mean")
# Project text embeddings to match image encoder dimension
text_projection = Dense(text_encoder.get_embedding_dimension(), 768)

# SigLIP outputs pooled embeddings directly, so no separate Pooling module is needed
image_encoder = Transformer("google/siglip2-base-patch16-224")

# Route inputs based on modality
router = Router(
    sub_modules={
        "text": [text_encoder, text_pooling, text_projection],
        "image": [image_encoder],
    },
)

model = SentenceTransformer(modules=[router])

Router 根据输入类型选择路线:普通文本进入文本编码器,图像 URL 或 PIL 图片进入图像编码器。

text_embeddings = model.encode(["A photo of a cat", "A pollinator on a flower"])
image_embeddings = model.encode([
    "https://huggingface.co/datasets/huggingface/cats-image/resolve/main/cats_image.jpeg",
    "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/bee.jpg",
])
print(text_embeddings.shape, image_embeddings.shape)
# (2, 768) (2, 768)

similarity = model.similarity(text_embeddings, image_embeddings)
print(similarity)
# tensor([[ 0.0028, -0.0144],
#         [-0.0233, -0.0355]])

向量维数一致,并不意味着向量空间已经对齐。上面的相似度数值只是原文未对齐结构的示例输出,不能作为图文匹配质量的结论。不同模态的编码器需要联合训练,才能在共享空间中得到有意义的相似度。也可结合任务路由,例如区分查询与文档编码路径;高级映射设置见 Router 文档。

数据集:先匹配损失,再检查列顺序

现代 SentenceTransformerTrainer 使用 Hugging Face datasets.Dataset,或表示多个数据集的 DatasetDict、名称到 Dataset 的字典。

从 Hub 加载

from datasets import load_dataset

train_dataset = load_dataset("sentence-transformers/all-nli", "pair-class", split="train")
eval_dataset = load_dataset("sentence-transformers/all-nli", "pair-class", split="dev")

print(train_dataset)
"""
Dataset({
    features: ['premise', 'hypothesis', 'label'],
    num_rows: 942069
})
"""

示例中的 942069 行是原文给出的数据集描述,未在本机下载核验。有些数据集需要同时指定子集。AllNLI 就提供 pair、pair-class、pair-score 和 triplet 四种格式,分别适合不同目标;不能只换子集名称而忽略损失函数的输入要求。Hub 上标有 sentence-transformers 的数据集可作为候选,但仍需检查语言、许可、数据质量和任务是否相符。

从本地文件加载

常见 CSV、JSON、Parquet、Arrow 等格式可通过 Datasets 的相应加载方式读取;SQL 数据需使用其支持的读取接口。CSV 和 JSON 示例为:

from datasets import load_dataset

dataset = load_dataset("csv", data_files="my_file.csv")
from datasets import load_dataset

dataset = load_dataset("json", data_files="my_file.json")

需要清理、过滤或预处理的数据,可以先整理成等长列表,再生成 Dataset:

from datasets import Dataset

anchors = []
positives = []
# Open a file, do preprocessing, filtering, cleaning, etc.
# and append to the lists

dataset = Dataset.from_dict({
    "anchor": anchors,
    "positive": positives,
})

这里的空列表是待填充的占位结构,不能直接训练。每个字典键会变成一列;样本数量、缺失值和列内类型应在构造前确认。

标签列与输入列的规则

首先按所选损失的要求检查是否需要标签。名为 label、labels、score 或 scores 的列会被视为标签;剩余列会被视为输入。输入的数量必须匹配损失函数,而且顺序决定角色,列名并不替你纠正顺序。

例如两列文本加一个浮点相似度标签,可以用于 CoSENTLoss、AnglELoss 或 CosineSimilarityLoss。三元组损失则需要按 anchor、positive、negative 排列输入。如果列顺序是 good_answer, bad_answer, question,程序仍可能接受这三个输入,却会把好答案当 anchor、坏答案当 positive、问题当 negative,训练目标随之出错。

可用 Dataset.select_columns 显式选择和排列列,或用 remove_columns 删除 sample_id、metadata、source、type 等额外信息。保留这些非标签列会使它们也成为模型输入,不能假设 Trainer 自动忽略。

多模态数据

多模态模型的输入列还可以包含图片、音频、视频以及多模态字典。标签和输入列规则保持相同。文档截图检索数据可包含查询、正例图片和若干负例图片:

from datasets import load_dataset

dataset = load_dataset("tomaarsen/llamaindex-vdr-en-train-preprocessed", "train", split="train")
"""
Dataset({
    features: ['query', 'image', 'negative_0', 'negative_1', 'negative_2', 'negative_3'],
    num_rows: 10000
})
"""

原文给出的列名和行数用于说明结构,并非本文下载结果。模型的数据整理器通过 preprocess 处理多模态预处理,一般无须在训练脚本里重复实现 tokenization 或图像处理;输入类型仍须与模型支持的模态相符。

损失函数:用数据格式约束选择

损失衡量当前批次的模型表现,优化器据此更新权重。没有一种损失适合所有数据和任务:成对类别、连续相似度、正例配对与三元组分别提供不同监督信息。

多数损失以待训练模型为主要参数。下面的 CoSENTLoss 需要两段文本和一个浮点相似度标签,因此选择 AllNLI 的 pair-score 子集:

from datasets import load_dataset
from sentence_transformers import SentenceTransformer
from sentence_transformers.sentence_transformer.losses import CoSENTLoss

# Load a model to train/finetune
model = SentenceTransformer("FacebookAI/xlm-roberta-base", model_kwargs={"torch_dtype": "float32"})

# Initialize the CoSENTLoss
# This loss requires pairs of text and a float similarity score as a label
loss = CoSENTLoss(model)
# Load an example training dataset that works with our loss function:
train_dataset = load_dataset("sentence-transformers/all-nli", "pair-score", split="train")
"""
Dataset({
    features: ['sentence1', 'sentence2', 'score'],
    num_rows: 942069
})
"""

更多选择应对照 Loss Overview 的输入格式要求。标签的范围和含义也须与具体损失定义一致,不能把类别编号直接当连续相似度。

训练参数:同时控制训练和观察过程

SentenceTransformerTrainingArguments 可配置优化、批次、精度以及评估和保存行为。虽然是可选组件,实际训练宜显式设置,以便记录和复现。

目的 常用参数
优化与训练长度 learning_rate、lr_scheduler_type、warmup_steps、num_train_epochs、max_steps、optim
批次与显存 per_device_train_batch_size、per_device_eval_batch_size、auto_find_batch_size、gradient_accumulation_steps、gradient_checkpointing、eval_accumulation_steps
精度与最佳模型 fp16、bf16、load_best_model_at_end、metric_for_best_model
输入组织 batch_sampler、multi_dataset_batch_sampler、prompts、router_mapping、learning_rate_mapping
评估与检查点 eval_strategy、eval_steps、save_strategy、save_steps、save_total_limit
日志与外部记录 report_to、run_name、log_level、logging_steps
Hub 保存 push_to_hub、hub_model_id、hub_strategy、hub_private_repo

下面保留当前官方页面的完整配置。FP16 和 BF16 须按硬件支持选择,不能同时照抄成适用于任何设备的设置。文档里的小数步数参数表示按训练步数比例设置相应间隔或预热;具体接受范围与语义需对应所安装的 Sentence Transformers/Transformers 版本。

args = SentenceTransformerTrainingArguments(
    # Required parameter:
    output_dir="models/mpnet-base-all-nli-triplet",
    # Optional training parameters:
    num_train_epochs=1,
    per_device_train_batch_size=16,
    per_device_eval_batch_size=16,
    learning_rate=2e-5,
    warmup_steps=0.1,
    fp16=True,  # Set to False if you get an error that your GPU can't run on FP16
    bf16=False,  # Set to True if you have a GPU that supports BF16
    batch_sampler=BatchSamplers.NO_DUPLICATES,  # losses that use "in-batch negatives" benefit from no duplicates
    # Optional tracking/debugging parameters:
    eval_strategy="steps",
    eval_steps=0.1,
    save_strategy="steps",
    save_steps=0.1,
    save_total_limit=2,
    logging_steps=0.01,
    run_name="mpnet-base-all-nli-triplet",  # Will be used in W&B if `wandb` is installed
)

NO_DUPLICATES 批采样有利于使用批内负例的损失,减少重复样本造成的负例冲突。它并不能代替数据本身的语义去重。原文单独展示的参数片段依赖相关类的导入;下文完整 Trainer 脚本包含这些导入。

评估器:把损失变成任务指标

eval_dataset 可用于训练期间计算评估损失;评估器则提供更贴近任务的指标。两者可以同时配置,也可以只使用其中一种。评估时机由 eval_strategy 和 eval_steps 控制。

评估器 所需数据
BinaryClassificationEvaluator 文本对和类别标签
EmbeddingSimilarityEvaluator 文本对和相似度分数
InformationRetrievalEvaluator 查询 ID 到查询、文档 ID 到文档、查询 ID 到相关文档 ID 集合
NanoBEIREvaluator 不要求手动传入数据,会自行加载基准
MSEEvaluator 由教师编码的源句和由学生编码的目标句,两组可相同
ParaphraseMiningEvaluator 句子 ID 映射与重复句子 ID 对
RerankingEvaluator 包含 query、positive 列表、negative 列表的样本字典
TranslationEvaluator 两种语言中的句子对
TripletEvaluator anchor、positive、negative 三元组

需要组合多个评估器时,可使用 SequentialEvaluator。

以 STSb 的连续相似度任务为例:

from datasets import load_dataset
from sentence_transformers.sentence_transformer.evaluation import EmbeddingSimilarityEvaluator, SimilarityFunction

# Load the STSB dataset (https://huggingface.co/datasets/sentence-transformers/stsb)
eval_dataset = load_dataset("sentence-transformers/stsb", split="validation")

# Initialize the evaluator
dev_evaluator = EmbeddingSimilarityEvaluator(
    sentences1=eval_dataset["sentence1"],
    sentences2=eval_dataset["sentence2"],
    scores=eval_dataset["score"],
    main_similarity=SimilarityFunction.COSINE,
    name="sts-dev",
)
# You can run evaluation like so:
# results = dev_evaluator(model)

AllNLI 三元组评估可只取一千个样本,以控制频繁评估的成本:

from datasets import load_dataset
from sentence_transformers.sentence_transformer.evaluation import TripletEvaluator, SimilarityFunction

# Load triplets from the AllNLI dataset (https://huggingface.co/datasets/sentence-transformers/all-nli)
max_samples = 1000
eval_dataset = load_dataset("sentence-transformers/all-nli", "triplet", split=f"dev[:{max_samples}]")

# Initialize the evaluator
dev_evaluator = TripletEvaluator(
    anchors=eval_dataset["anchor"],
    positives=eval_dataset["positive"],
    negatives=eval_dataset["negative"],
    main_similarity_function=SimilarityFunction.COSINE,
    name="all-nli-dev",
)
# You can run evaluation like so:
# results = dev_evaluator(model)

没有自行准备的评估数据时,也可使用自动加载基准的 NanoBEIR:

from sentence_transformers.sentence_transformer.evaluation import NanoBEIREvaluator

# Initialize the evaluator. Unlike most other evaluators, this one loads the relevant datasets
# directly from Hugging Face, so there's no mandatory arguments
dev_evaluator = NanoBEIREvaluator()
# You can run evaluation like so:
# results = dev_evaluator(model)

自动加载仍会涉及网络和缓存,并非无数据或无外部依赖。频繁评估时,较小的开发集能减少开销。原文提出 90% 训练、1% 开发、9% 测试的可选划分,但比例应依据数据量和任务调整。测试集应留给最终评估,避免参与持续调参。

训练完成后可用 trainer.evaluate(test_dataset) 查看测试损失,或用独立 test evaluator 计算任务指标。如果在保存模型前执行评估,自动生成的模型卡可记录这些结果。分布式训练中,评估器只在第一个设备运行;训练和评估数据集则由多个设备共同处理,二者的运行方式不同。

Trainer:完整训练例子

Trainer 汇集模型、训练数据、损失、可选参数、可选开发集和评估器。下面完整保留 MPNet 在 AllNLI 三元组上训练的官方例子,包括基线评估、训练、测试评估和模型保存。

from datasets import load_dataset
from sentence_transformers import (
    SentenceTransformer,
    SentenceTransformerTrainer,
    SentenceTransformerTrainingArguments,
    SentenceTransformerModelCardData,
)
from sentence_transformers.sentence_transformer.losses import MultipleNegativesRankingLoss
from sentence_transformers.sentence_transformer.training_args import BatchSamplers
from sentence_transformers.sentence_transformer.evaluation import TripletEvaluator

# 1. Load a model to finetune with 2. (Optional) model card data
model = SentenceTransformer(
    "microsoft/mpnet-base",
    model_card_data=SentenceTransformerModelCardData(
        language="en",
        license="apache-2.0",
        model_name="MPNet base trained on AllNLI triplets",
    ),
    model_kwargs={"torch_dtype": "float32"},
)
# 3. Load a dataset to finetune on
dataset = load_dataset("sentence-transformers/all-nli", "triplet")
train_dataset = dataset["train"].select(range(100_000))
eval_dataset = dataset["dev"]
test_dataset = dataset["test"]

# 4. Define a loss function
loss = MultipleNegativesRankingLoss(model)

# 5. (Optional) Specify training arguments
args = SentenceTransformerTrainingArguments(
    # Required parameter:
    output_dir="models/mpnet-base-all-nli-triplet",
    # Optional training parameters:
    num_train_epochs=1,
    per_device_train_batch_size=16,
    per_device_eval_batch_size=16,
    learning_rate=2e-5,
    warmup_steps=0.1,
    fp16=True,  # Set to False if you get an error that your GPU can't run on FP16
    bf16=False,  # Set to True if you have a GPU that supports BF16
    batch_sampler=BatchSamplers.NO_DUPLICATES,  # MultipleNegativesRankingLoss benefits from no duplicate samples in a batch
    # Optional tracking/debugging parameters:
    eval_strategy="steps",
    eval_steps=0.1,
    save_strategy="steps",
    save_steps=0.1,
    save_total_limit=2,
    logging_steps=0.01,
    run_name="mpnet-base-all-nli-triplet",  # Will be used in W&B if `wandb` is installed
)

# 6. (Optional) Create an evaluator & evaluate the base model
dev_evaluator = TripletEvaluator(
    anchors=eval_dataset["anchor"],
    positives=eval_dataset["positive"],
    negatives=eval_dataset["negative"],
    name="all-nli-dev",
)
dev_evaluator(model)

# 7. Create a trainer & train
trainer = SentenceTransformerTrainer(
    model=model,
    args=args,
    train_dataset=train_dataset,
    eval_dataset=eval_dataset,
    loss=loss,
    evaluator=dev_evaluator,
)
trainer.train()

# (Optional) Evaluate the trained model on the test set
test_evaluator = TripletEvaluator(
    anchors=test_dataset["anchor"],
    positives=test_dataset["positive"],
    negatives=test_dataset["negative"],
    name="all-nli-test",
)
test_evaluator(model)

# 8. Save the trained model
model.save_pretrained("models/mpnet-base-all-nli-triplet/final")

# 9. (Optional) Push it to the Hugging Face Hub
model.push_to_hub("mpnet-base-all-nli-triplet")

该脚本会读取远程模型和数据,若执行最后一步还会将模型推到 Hugging Face Hub。此处只保留完整示例结构,没有下载、训练、评估或上传。实际使用时,应先配置环境和数据许可,并按自己的需要删除或关闭外部上传与日志步骤。

SentenceTransformerModelCardData 中的语言、许可和名称描述的是示例模型卡,不能把这些值自动套到任意训练数据和基座的组合。当前页面的导入路径包含 sentence_transformers.sentence_transformer 子命名空间;旧版本可能使用不同路径,应对照所安装版本,而不要将新旧代码混拼。

回调

Trainer 支持 Transformers 的 TrainerCallback 子类。安装并启用相应组件后,WandbCallback 可记录 W&B 指标,TensorBoardCallback 可记录 TensorBoard 日志,CodeCarbonCallback 可估算训练碳排放并写入模型卡。也可编写自己的回调。本文没有连接这些服务,也没有产生训练日志。

多数据集训练:让格式与损失按名称对应

多任务训练可以保留每个 Dataset 的不同格式,无须先强行转换成一种样本结构。给 train_dataset 传入名称到 Dataset 的字典;需要不同损失时,再传入使用相同名称的损失字典。eval_dataset 同样可使用字典,但不要求每个训练集都对应一个评估集。

每个批次只来自一个数据集。跨数据集的取样由 multi_dataset_batch_sampler 决定:

  • ROUND_ROBIN 轮流取各数据集的批次,任意一个数据集耗尽就结束。因此取样机会相同,但部分大数据集样本可能没有用到。
  • 默认 PROPORTIONAL 按大小比例取样,较大数据集出现更频繁,并覆盖各数据集的样本。

下面完整示例包括 AllNLI 的四种子集、STSb、Quora 重复问题和 Natural Questions,并分别配置 MultipleNegativesRankingLoss、SoftmaxLoss 和 CoSENTLoss。

from datasets import load_dataset
from sentence_transformers import SentenceTransformer, SentenceTransformerTrainer
from sentence_transformers.sentence_transformer.losses import CoSENTLoss, MultipleNegativesRankingLoss, SoftmaxLoss

# 1. Load a model to finetune
model = SentenceTransformer("google-bert/bert-base-uncased", model_kwargs={"torch_dtype": "float32"})

# 2. Load several Datasets to train with
# (anchor, positive)
all_nli_pair_train = load_dataset("sentence-transformers/all-nli", "pair", split="train[:10000]")
# (premise, hypothesis) + label
all_nli_pair_class_train = load_dataset("sentence-transformers/all-nli", "pair-class", split="train[:10000]")
# (sentence1, sentence2) + score
all_nli_pair_score_train = load_dataset("sentence-transformers/all-nli", "pair-score", split="train[:10000]")
# (anchor, positive, negative)
all_nli_triplet_train = load_dataset("sentence-transformers/all-nli", "triplet", split="train[:10000]")
# (sentence1, sentence2) + score
stsb_pair_score_train = load_dataset("sentence-transformers/stsb", split="train[:10000]")
# (anchor, positive)
quora_pair_train = load_dataset("sentence-transformers/quora-duplicates", "pair", split="train[:10000]")
# (query, answer)
natural_questions_train = load_dataset("sentence-transformers/natural-questions", split="train[:10000]")
# We can combine all datasets into a dictionary with dataset names to datasets
train_dataset = {
    "all-nli-pair": all_nli_pair_train,
    "all-nli-pair-class": all_nli_pair_class_train,
    "all-nli-pair-score": all_nli_pair_score_train,
    "all-nli-triplet": all_nli_triplet_train,
    "stsb": stsb_pair_score_train,
    "quora": quora_pair_train,
    "natural-questions": natural_questions_train,
}

# 3. Load several Datasets to evaluate with
# (anchor, positive, negative)
all_nli_triplet_dev = load_dataset("sentence-transformers/all-nli", "triplet", split="dev")
# (sentence1, sentence2, score)
stsb_pair_score_dev = load_dataset("sentence-transformers/stsb", split="validation")
# (anchor, positive)
quora_pair_dev = load_dataset("sentence-transformers/quora-duplicates", "pair", split="train[10000:11000]")
# (query, answer)
natural_questions_dev = load_dataset("sentence-transformers/natural-questions", split="train[10000:11000]")

# We can use a dictionary for the evaluation dataset too, but we don't have to. We could also just use
# no evaluation dataset, or one dataset.
eval_dataset = {
    "all-nli-triplet": all_nli_triplet_dev,
    "stsb": stsb_pair_score_dev,
    "quora": quora_pair_dev,
    "natural-questions": natural_questions_dev,
}
# 4. Load several loss functions to train with
# (anchor, positive), (anchor, positive, negative)
mnrl_loss = MultipleNegativesRankingLoss(model)
# (sentence_A, sentence_B) + class
softmax_loss = SoftmaxLoss(model, model.get_embedding_dimension(), 3)
# (sentence_A, sentence_B) + score
cosent_loss = CoSENTLoss(model)

# Create a mapping with dataset names to loss functions, so the trainer knows which loss to apply where.
# Note that you can also just use one loss if all of your training/evaluation datasets use the same loss
losses = {
    "all-nli-pair": mnrl_loss,
    "all-nli-pair-class": softmax_loss,
    "all-nli-pair-score": cosent_loss,
    "all-nli-triplet": mnrl_loss,
    "stsb": cosent_loss,
    "quora": mnrl_loss,
    "natural-questions": mnrl_loss,
}

# 5. Define a simple trainer, although it's recommended to use one with args & evaluators
trainer = SentenceTransformerTrainer(
    model=model,
    train_dataset=train_dataset,
    eval_dataset=eval_dataset,
    loss=losses,
)
trainer.train()

# 6. save the trained model and optionally push it to the Hugging Face Hub
model.save_pretrained("bert-base-all-nli-stsb-quora-nq")
model.push_to_hub("bert-base-all-nli-stsb-quora-nq")

其中 Quora 与 Natural Questions 的开发片段取自训练分割的后续区间,和前一万个训练样本不重叠;正式实验仍应检查更深层的重复或内容泄漏。SoftmaxLoss 的类别数 3 对应该示例的 NLI 分类,不能原封不动用于类别数量不同的数据。

官方页面引用了 Huang 等人在中文任务中结合不同损失以及 MatryoshkaLoss 的研究结果,说明多任务训练和可变长度嵌入的用途。这是原文引用的外部研究,不是本文新训练出的结果;具体实现仍应阅读所对应论文和训练脚本。

旧训练方法与兼容边界

Sentence Transformers v3.0 之前,常见流程使用 InputExample、PyTorch DataLoader 和 SentenceTransformer.fit():

from sentence_transformers import SentenceTransformer, InputExample, losses
from torch.utils.data import DataLoader

# Define the model. Either from scratch of by loading a pre-trained model
model = SentenceTransformer("distilbert/distilbert-base-uncased")
# Define your train examples. You need more than just two examples...
train_examples = [
    InputExample(texts=["My first sentence", "My second sentence"], label=0.8),
    InputExample(texts=["Another pair", "Unrelated sentence"], label=0.3),
]

# Define your train dataset, the dataloader and the train loss
train_dataloader = DataLoader(train_examples, shuffle=True, batch_size=16)
train_loss = losses.CosineSimilarityLoss(model)

# Tune the model
model.fit(train_objectives=[(train_dataloader, train_loss)], epochs=1, warmup_steps=100)

两个样本仅用于演示格式,不足以支撑有意义的训练。从 v3.0 开始,fit() 仍可供既有脚本使用,但内部会构造 Trainer;直接使用 Trainer 能更完整地控制训练参数。若更新后的 fit() 与旧脚本不兼容,可参考 old_fit() 恢复旧行为,官方同时说明该方法计划在未来完全弃用。

新项目应优先使用现代 Trainer。迁移旧项目时,要先辨认它依赖的实际版本、输入格式和损失,不能仅把函数名替换后就声称结果等价。

基座模型比较的适用范围

原文使用一项固定实验比较基座:在 56 万组三元组上训练一个 epoch,批次为 64,再评估 14 项不同领域的句子相似度任务。下面完整保留该实验结果。它是原文的基准记录,既不是本文实测,也不是当前所有嵌入模型的通用排名。

模型 14 项任务的表现
microsoft/mpnet-base 60.99
nghuyong/ernie-2.0-en 60.73
microsoft/deberta-base 60.21
FacebookAI/roberta-base 59.63
google-t5/t5-base 59.21
google-bert/bert-base-uncased 59.17
distilbert/distilbert-base-uncased 59.03
nreimers/TinyBERT_L-6_H-768_v2 58.27
google/t5-v1_1-base 57.63
nreimers/MiniLMv2-L6-H768-distilled-from-BERT-Large 57.31
albert/albert-base-v2 57.14
microsoft/MiniLM-L12-H384-uncased 56.79
microsoft/deberta-v3-base 54.46

GLUE 或 SuperGLUE 上的更高分数不自动意味着更好的句子嵌入。目标语言、训练监督和真实使用场景仍决定哪个模型适用。

与 CrossEncoder 训练的区别

两种训练流程相似,但输入限制不同:CrossEncoder 训练允许一列中保存长度可变的文本列表;SentenceTransformer 的上述训练方式不允许一列装入可变数量的文本输入。因此不能据此构造每个样本负例数量不同的列表,再假定 Trainer 会自动适配。具体 CrossEncoder 流程见 CrossEncoder Training Overview。

完整的文本嵌入端到端训练与微调、多模态嵌入和重排模型示例,见原页面末尾的 Hugging Face 教程链接。它们展示了从基座、训练到对照基准的全过程,本文没有执行这些外部教程。

官方页面还给 AI 编码代理提供可选技能安装命令:

hf skills add train-sentence-transformers [--claude] [--global]

方括号表示文档里的可选参数提示,不能整段连同括号机械执行。本文没有安装该技能或改变代理配置。

来源、许可与改动

来源:Training Overview — Sentence Transformers。作者与项目:Nils Reimers 及 Sentence Transformers 项目贡献者,Copyright 2019 Nils Reimers。源仓库以 Apache License 2.0 发布。

本文对官方技术正文作完整中文翻译与编辑核对,保留模型类型、训练组件、主要章节、23 个代码块、评估器与基准表;导航、折叠标签和网页 UI 转为顺序章节。代码保留当前在线页面内容,增加了版本、占位数据、未训练结果与外部上传的说明。所有示例未执行,未下载模型和数据,未测量性能。许可全文同时保存在随稿资产中。

Apache License 2.0 完整许可及原版权声明
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/

TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION

1. Definitions.

“License” shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.

“Licensor” shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
“Legal Entity” shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
“control” means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
“You” (or “Your”) shall mean an individual or Legal Entity
exercising permissions granted by this License.

“Source” form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
“Object” form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.

“Work” shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
“Derivative Works” shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
“Contribution” shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, “submitted”
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as “Not a Contribution.”
“Contributor” shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:

(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and

(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a “NOTICE” text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an “AS IS” BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS

APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets “{}”
replaced with your own identifying information. (Don’t include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same “printed page” as the copyright notice for easier
identification within third-party archives.
Copyright 2019 Nils Reimers

Licensed under the Apache License, Version 2.0 (the “License”);
you may not use this file except in compliance with the License.
You may obtain a copy of the License at

http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an “AS IS” BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.

© 版权声明
THE END
喜欢就支持一下吧
点赞0 分享
评论 抢沙发

请登录后发表评论

    暂无评论内容