Hugging Face Sentence Transformers 训练指南(历史参考)

TL;DR

Hugging Face 发布了一篇历史指南,详细讲解了构建、训练和微调 Sentence Transformers 模型的全过程,涵盖了模型架构、数据集准备、损失函数选择以及模型发布,但指出文中使用的 SentenceTransformer.fit API 已废弃,并引导读者参考基于 SentenceTransformerTrainer 的新指南。


指南概览

该教程仅作参考保留;它说明了如何从零创建 Sentence Transformers 模型或微调已有模型,如何组织训练数据、哪种损失函数对应哪种数据格式,以及如何将生成的模型推送至 Hugging Face Hub。

注意: 本指南使用的是 v3.0 之前的 SentenceTransformer.fit API,已被 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 将可变长度的文本(或图像)映射为固定维度的嵌入向量,以捕获语义含义。

  1. Transformer 层 – 输入文本由预训练的 Transformer(例如 distilroberta-base)处理,模型输出带上下文的 token 嵌入。
  2. 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 的新版训练指南,以获得生产就绪的工作流。

Sources