数据截至 (上游 commit 5cc889d47547)
Sentence Transformers — 架构与原理
30 秒导读: Sentence Transformers(通称 SBERT,因 2019 年论文 Sentence-BERT 得名)是把文本变成「语义向量」的标准库:句子进去、向量出来,向量之间距离近 = 语义近。它同时是推理库(
encode一键取向量)和训练库(用对比学习微调出你自己的嵌入模型),并在本 commit(v6.1.0.dev0)已扩展成四个编码器家族:稠密双塔SentenceTransformer、逐对打分的CrossEncoder、词项稀疏的SparseEncoder(SPLADE)、token 级多向量的MultiVectorEncoder(ColBERT)。
1. 这是什么(零基础也能懂)
一句话定义
一个把句子/段落编码成定长向量、并让「向量距离 ≈ 语义距离」的模型训练与推理框架。
它解决谁的什么问题
BERT 这类模型直接输出的是每个 token 一个向量,而且原生向量空间里「语义相似的两个句子」并不靠近。想做「搜出和这句话意思相近的一百万条文本」时,逐对过一遍 BERT 要算 5 千万次、以小时计(SBERT 论文的经典账)。
Sentence Transformers 的答案是双塔(bi-encoder):用一个共享权重的模型分别把每句话压成一个向量,离线把全库编码完,线上一次查询只是一次向量近邻搜索——毫秒级。
它能做什么
| 能力 | 入口 |
|---|---|
| 稠密句嵌入(语义相似、聚类、语义搜索) | SentenceTransformer.encode(sentence_transformers/sentence_transformer/model.py:754) |
| 逐对精准打分 / 重排序(rerank) | CrossEncoder.predict / CrossEncoder.rank(sentence_transformers/cross_encoder/model.py:752) |
| 可倒排索引的稀疏向量(词汇+扩展词项) | SparseEncoder(sentence_transformers/sparse_encoder/model.py:34) |
| token 级多向量迟交互检索 | MultiVectorEncoder(sentence_transformers/multi_vector_encoder/model.py:65) |
| 训练以上全部模型 | SentenceTransformerTrainer 等(sentence_transformers/sentence_transformer/trainer.py:36) |
| 难负例挖掘、量化、ONNX/OpenVINO 导出 | mine_hard_negatives、quantize_embeddings、sentence_transformers/backend/ |
用起来什么样
最小示例(摘自 SentenceTransformer 文档字符串,sentence_transformers/sentence_transformer/model.py:116-139):
from sentence_transformers import SentenceTransformer
model = SentenceTransformer("sentence-transformers/all-mpnet-base-v2")
embeddings = model.encode(["今天天气真好", "外面阳光灿烂", "他开车去了体育场"])
print(embeddings.shape) # (3, 768)
similarities = model.similarity(embeddings, embeddings)
# 前两句余弦相似度 ≈0.68,与第三句 ≈0.05 —— 距离就是语义
训练侧同样只要五行核心代码:准备一个「列=文本」的 datasets.Dataset,选一个损失,交给 trainer:
# 示意,非源码(真实版见 examples/sentence_transformer/training/ms_marco/train_bi_encoder_mnrl.py)
model = SentenceTransformer("microsoft/mpnet-base")
loss = MultipleNegativesRankingLoss(model) # 损失在构造时拿到模型
trainer = SentenceTransformerTrainer(model=model, train_dataset=train_dataset, loss=loss)
trainer.train()