Hugging Face Sentence Transformers 训练指南(历史参考)
TL;DR
Hugging Face 发布了一篇历史指南,详细讲解了构建、训练和微调 Sentence Transformers 模型的全过程,涵盖了模型架构、数据集准备、损失函数选择以及模型发布,但指出文中使用的 SentenceTransformer.fit API 已废弃,并引导读者参考基于 SentenceTransformerTrainer 的新指南。
指南概览
该教程仅作参考保留;它说明了如何从零创建 Sentence Transformers 模型或微调已有模型,如何组织训练数据、哪种损失函数对应哪种数据格式,以及如何将生成的模型推送至 Hugging Face Hub。
注意: 本指南使用的是 v3.0 之前的
SentenceTransformer.fitAPI,已被SentenceTransformerTrainer取代。当前的训练流程记录在以下最新文章中:
- 嵌入模型 – Training and Finetuning Embedding Models with Sentence Transformers
- 重新排序模型 – Training and Finetuning Reranker Models with Sentence Transformers
- 稀疏嵌入模型 – Training and Finetuning Sparse Embedding Models with Sentence Transformers
- 多模态模型 – Training and Finetuning Multimodal Embedding & Reranker Models with Sentence Transformers
Sentence Transformers 模型工作原理
Sentence Transformers 将可变长度的文本(或图像)映射为固定维度的嵌入向量,以捕获语义含义。
- Transformer 层 – 输入文本由预训练的 Transformer(例如
distilroberta-base)处理,模型输出带上下文的 token 嵌入。 - Pooling 层 – 将 token 嵌入聚合(例如均值池化)为单个句子级向量。
from sentence_transformers import SentenceTransformer, models
# Layer 1: pre‑trained transformer
word_embedding_model = models.Transformer('distilroberta-base')
# Layer 2: pooling to a fixed‑size vector
pooling_model = models.Pooling(word_embedding_model.get_word_embedding_dimension())
# Assemble the modules
model = SentenceTransformer(modules=[word_embedding_model, pooling_model])
模型是一个按顺序排列的模块列表;如有需要,可插入额外的层(全连接、卷积等)。
为什么不直接使用普通的 Transformer 来生成句子嵌入?
- 在 10,000 条句子上使用原始 BERT 模型进行语义搜索的推理需要约 5 千万次运算(约 65 小时),而 Sentence Transformer 将其降至约 5 秒。
- 直接对 BERT token 嵌入取平均会得到比经典 GloVe 嵌入更差的句子表示。
准备数据集
训练需要提供句子相似或不相似的信号。指南列出了四种常见的数据集结构:
| 案例 | 格式 | 常见来源 | 推荐损失 |
|---|---|---|---|
| 1 | (sentence_a, sentence_b, similarity_label) – 标签可以是整数或浮点数 |
自然语言推理(NLI)数据集 | ContrastiveLoss, SoftmaxLoss, CosineSimilarityLoss |
| 2 | (sentence_a, sentence_b) – 正向对,无显式标签 |
同义句、摘要、重复提问对 | MultipleNegativesRankingLoss, MegaBatchMarginLoss |
| 3 | (sentence, class_id) – 整数类别标签 |
主题分类数据集(例如 TREC) | 使用类别 ID 的基于三元组的损失(如 BatchHardTripletLoss 等) |
| 4 | (anchor, positive, negative) – 显式三元组,无类别 ID |
预构建的三元组数据集(例如 Quora Triplets) | TripletLoss |
教程演示了使用 embedding-data/QQP_triplets 数据集的案例 4。它展示了如何使用 datasets.load_dataset 加载数据集、检查其结构,并将每个示例转换为 sentence_transformers.InputExample:
from datasets import load_dataset
from sentence_transformers import InputExample
dataset = load_dataset('embedding-data/QQP_triplets')
train_examples = []
train_data = dataset['train']['set']
for i in range(dataset['train'].num_rows // 2): # use half the data for speed
ex = train_data[i]
train_examples.append(
InputExample(texts=[ex['query'], ex['pos'][0], ex['neg'][0]])
)
随后将这些示例包装进 torch.utils.data.DataLoader 以进行批处理:
from torch.utils.data import DataLoader
train_dataloader = DataLoader(train_examples, shuffle=True, batch_size=16)
选择损失函数
损失函数必须与数据集格式匹配:
- 案例 1 – 使用
ContrastiveLoss(整数标签)或CosineSimilarityLoss(浮点标签)。 - 案例 2 – 使用
MultipleNegativesRankingLoss(最常用)或MegaBatchMarginLoss。 - 案例 3 – 使用依赖类别 ID 的基于三元组的损失,例如
BatchHardTripletLoss。 - 案例 4 – 使用
TripletLoss,它不需要类别标签。
实例化损失函数的代码非常简洁:
from sentence_transformers import losses
train_loss = losses.TripletLoss(model=model)
训练 / 微调模型
准备好 DataLoader 和损失后,训练只需一次 fit 调用即可进行:
model.fit(train_objectives=[(train_dataloader, train_loss)], epochs=10)
如果微调已有模型(例如 sentence-transformers/all-MiniLM-L6-v2),可通过 SentenceTransformer(model_id) 加载模型后直接调用 fit。
将模型发布到 Hub
训练完成后,将模型推送至 Hugging Face Hub:
from huggingface_hub import notebook_login
notebook_login() # or `huggingface-cli login` in a terminal
model.save_to_hub(
"distilroberta-base-sentence-transformer",
organization="<your‑username-or‑org>",
train_datasets=["embedding-data/QQP_triplets"]
)
save_to_hub 会自动生成模型卡片、推理小部件以及示例代码片段。
Sentence Transformers 的局限性
Sentence Transformers 在语义搜索和相似度任务上表现出色,但不适用于纯分类问题。进行分类时,应使用标准的 🤗 Transformers 库(例如序列分类 pipeline)。
其他资源
Getting Started With Embeddings – 嵌入入门指南。
Understanding Semantic Search – 深入了解语义检索。
Your First Sentence Transformers Model – 步骤详尽的入门教程。
Playlist Generator – Sentence Transformers 的示例应用。
Hugging Face + Sentence Transformers documentation – 完整的 API 参考。
本指南仅作历史参考保留;请查阅使用
SentenceTransformerTrainer的新版训练指南,以获得生产就绪的工作流。