Hugging Face Sentence Transformers 訓練指南(歷史參考)

TL;DR

Hugging Face 發布了一篇歷史性的指南,說明如何構建、訓練與微調 Sentence Transformers 模型,涵蓋架構、資料集準備、損失函式選擇與模型發布,但指出文中描述的 SentenceTransformer.fit API 已過時,並引導讀者參考較新的基於 SentenceTransformerTrainer 的指南。


指南概覽

此教學僅保留作為參考;它說明如何從頭建立 Sentence Transformers 模型或微調現有模型、如何格式化訓練資料、哪種損失函式對應每種格式,以及如何將產生的模型推送至 Hugging Face Hub。

注意: 本指南使用 pre‑v3.0 SentenceTransformer.fit API,已被 SentenceTransformerTrainer 取代。最新的訓練流程記錄於以下最新文章:

  • 嵌入模型 – Training and Finetuning Embedding Models with Sentence Transformers
  • Reranker 模型 – 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,000 萬次運算(約 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