跳到主要内容

数据截至 (上游 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_negativesquantize_embeddingssentence_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()

一句话直觉

把这个库想成一条「文本流水线 + 对比教练」:流水线(Transformer → Pooling → …)负责把不定长的字句压成一枚向量;教练(对比损失)在训练时不停地说「这两个该近一点、那两个该远一点」,直到向量空间的远近真的等于语义的远近。


2. 顶层全景(它大概怎么转)

2.1 四个家族,一副骨架

本 commit 的包结构按「编码器家族」分目录,全部家族共享同一个基类 BaseModel——它直接继承 nn.Sequential(sentence_transformers/base/model.py:57):

┌─────────────────────────────────────┐
│ BaseModel(nn.Sequential) + Trainer │
│ base/model.py · base/trainer.py │
└───────┬─────────┬─────────┬─────────┘
│ │ │
┌───────────────────┘ │ └───────────────────┐
▼ ▼ ▼
SentenceTransformer CrossEncoder SparseEncoder /
(稠密双塔,出向量) (句对→分数,不出向量) MultiVectorEncoder
Transformer→Pooling 联合编码句对 (SPLADE / ColBERT)
→Normalize… →分类/回归头 →词项/token 向量
│ │ │
▼ ▼ ▼
离线建库+向量近邻 在线 rerank/STS 倒排索引 / MaxSim 检索

怎么读这张图: 中间是共享底座(模块链 + 训练器),下面四个家族的区别只在「模块链里装什么」和「输出是什么形态」。学一个等于学了四个的骨架。

2.2 部件一句话职责

部件干什么在哪个文件
BaseModel四族共同的 nn.Sequential 基类:加载、保存、设备、prompt 管理sentence_transformers/base/model.py:57
Transformer包一层 HF AutoModel:预处理(tokenize)→ 前向 → 写 token_embeddingssentence_transformers/base/modules/transformer.py:644
Pooling把变长 token 向量压成定长句向量(mean/cls/max/lasttoken…)sentence_transformers/sentence_transformer/modules/pooling.py:70
Dense / Normalize线性投影 / L2 归一化,链路末端的可选模块sentence_transformers/base/modules/dense.py:21normalize.py:14
Router按 task( query/document )或模态把输入路由到不同子链sentence_transformers/base/modules/router.py:29
BaseTrainer继承 HF Trainer,负责 collator 接线、loss 注入、batch samplersentence_transformers/base/trainer.py:76
BaseDataCollator逐列调 model.preprocess,产出 列名_input_ids 式扁平 batchsentence_transformers/base/data_collator.py:18
损失族MultipleNegativesRankingLossCoSENTLossTripletLoss 等 28 种sentence_transformers/sentence_transformer/losses/
batch sampler 族无重复/按标签分组/多数据集混合的采样策略sentence_transformers/base/sampler.py:35
mine_hard_negatives用已有模型检索「像锚点但不是」的文本当难负例sentence_transformers/util/hard_negatives.py:25

2.3 主线走一遍(推理:encode)

SentenceTransformer.encode(sentence_transformers/sentence_transformer/model.py:754)的一次调用:

① 输入文本列表


② 按长度降序排序(减少 padding 浪费) model.py:925


③ preprocess:拼 prompt → tokenize → features 字典 base/model.py:587
│ features = {input_ids, attention_mask, ...}

④ 沿模块链前向:Transformer → Pooling → (Dense/Normalize) base/model.py:554
│ 每个模块吃掉 features 字典、往里写新键、传给下一个

⑤ 取 features["sentence_embedding"] → 截断/归一/量化 → 恢复原始顺序 → numpy

要点只有一个:模块之间传的不是张量,而是一个不断累加的 features 字典——这是整个库的插件协议(第 1 章细讲)。

2.4 主线走一遍(训练:一个 step)

① DataLoader 按 batch sampler 取出一批样本(每行 = 若干文本列 + 可选 label)


② collator 对每一列分别 preprocess,列名做前缀拍平 base/data_collator.py:131
│ {anchor_input_ids..., positive_input_ids..., label}

③ collect_features 按前缀拆回 N 个 features 字典 base/trainer.py:584


④ loss(features, labels) —— 损失函数内部自己调用 model() 算嵌入
│ (锚/正/负各过一遍同一个模型,即「双塔共享权重」)

⑤ 对比损失:正例对拉近、batch 内其余全部当负例推远


⑥ HF Trainer 反向传播、更新

这条线最值得记住的两点:

  1. 损失函数持有模型。 每个 loss 构造时就拿到 model(sentence_transformers/sentence_transformer/losses/multiple_negatives_ranking.py:188),forward 里自己跑 embed_columns(self.model, sentence_features)——所以「怎么编码」与「怎么打分」住在同一个 nn.Module 里。
  2. batch 里的其他样本就是负例。 默认的 MNRL 不用单独准备负样本:一个 batch 的 B 个正例对,互为负例,等效每步见到 B−1 个负样本(第 2、3 章细讲)。

3. 阅读地图(建议顺序)

五章由浅入深。时间有限就读 01 → 02,这两章覆盖了「它为什么是标准库」的 80%。

顺序章节讲什么适合谁
101-module-pipeline.md模块链与 features 字典协议;encode 的长度排序、量化、多进程;modules.json 自描述存取所有人必读,这是全库的地基
202-contrastive-training.md双塔训练:loss 持有模型的接线、MNRL/InfoNCE 拆解、Matryoshka、GradCache、列合并前向想训练/微调嵌入模型的人
303-data-sampling.md数据集列序契约、四种 batch sampler、多数据集混合、难负例挖掘训练效果不对劲时来查这章
404-cross-encoder.mdCrossEncoder:联合编码为什么更准、predict/rank、排序损失、与双塔的分工做检索/重排序系统的人
505-sparse-and-multi-vector.mdSPLADE 稀疏向量与 ColBERT 多向量:同一副骨架上的另外两种输出形态关心可解释检索、倒排索引、迟交互的人

4. 巧妙之处(可借鉴的技术)

  1. 模块协议是「字典进、字典出」,不是张量。 每个模块读 features 里自己要的键、写入新键(sentence_transformers/base/modules/module.py:87Module.forward 抽象)。新增一种池化/投影只需新加一个模块类,链上其他成员零改动——modules.json 里登记的 type 就是完整的导入路径,保存即自描述(sentence_transformers/base/model.py:753)。
  2. 合并前向:三列文本一次过模型。 训练时锚/正/负三列本来要三次前向;merge_feature_batches 检查三列预处理结果同构后拼成一个大 batch 一次算完,再 reshape 切回来(sentence_transformers/base/losses/merged_forward.py:53)。同构判断失败就安静回退,保证正确性不受性能优化影响。
  3. GradCache:用「梯度缓存 + 重放」把对比学习的 batch 扩到显存装不下的规模。 先无梯度地逐 mini-batch 算嵌入并缓存,整批算 loss 得到嵌入的梯度,再带梯度重放每个 mini-batch 把缓存梯度注入(sentence_transformers/base/losses/gradcache.py:397forward_cached)。配合随机状态快照(RandContext,gradcache.py:31)保证重放时 dropout 一致。
  4. 难负例感知是有验证开关的。 MNRL 的 hardness_mode/hardness_strength 把「越像的负例惩罚越重」做成了一个 stop-grad 加分项,直接挂在 logits 上(multiple_negatives_ranking.py:286-320),对应 Lan et al. 2025 与 EmbeddingGemma 的做法。
  5. encode 的长度排序 + FA2 展平交错。 先按输入长度降序排序减少 padding;若走 flash-attention 展平(无 padding),则改用「最长、最短、次长、次短…」交错,让每个 batch 的 token 总数均匀、显存峰值平滑(sentence_transformers/base/model.py:487_interleave_sorted_indices)。
  6. 列序即语义的护栏。 数据集列的顺序决定谁是锚谁是正例;collator 发现 anchor 出现在非第 0 列等反常情况会主动警告(sentence_transformers/base/data_collator.py:150maybe_warn_about_column_order)——把一类最常见的训练事故变成了日志。

5. 边界与局限

  • 双塔的天花板是「交互太晚」。 query 和 document 各自独立编码,交叉注意力为零,所以精度上限低于 CrossEncoder(见第 4 章)。库的回答不是「双塔更好」,而是「两个都给你,用在对的环节」:双塔粗排、CrossEncoder 精排。
  • 向量只在同一模型内可比。 不同模型(甚至同一模型不同 checkpoint)产出的向量空间不可混用——换了模型必须全库重编码。这是嵌入模型的性质,不是本库的实现缺陷。
  • 榜单分数不等于你的领域表现。 卡片 aiRef/sources/sentence-transformers.md 的 Gotchas 原话:embedding 质量领域相关,MTEB 排名不迁移到你的数据。库内置了 InformationRetrievalEvaluator 等评估器(sentence_transformers/sentence_transformer/evaluation/information_retrieval.py:24)就是为了让你在自己的数据上量。
  • MNRL 对 batch 内重复敏感。 同一文本若同时以锚和「别人的正例」出现,会被当成负例推开——所以官方推荐配 BatchSamplers.NO_DUPLICATES(见第 3 章),但默认 sampler 并不开它。
  • encode 会重排输入再还原顺序(model.py:925:976)。对调用者透明,但意味着它是一次性收集式的批处理接口,不是流式的。
  • 稀疏/多向量是特化形态。 SPLADE 路线要求 MLM backbone(sparse_encoder/model.py:1108ForMaskedLM 架构探测),ColBERT 路线存储开销是每 token 一个向量——两者都不是「免费升级」,第 5 章有账。

6. 横向对比

同书架上的相邻 teardown:

  • transformers — 本库的 Transformer 模块就是 HF AutoModel 的一层壳(base/modules/transformer.py:644):transformers 提供「骨干网络与前向」,sentence-transformers 提供「把骨干变成嵌入模型的池化/训练/检索语义」。读骨干加载细节去那边,读「骨干之上」来这边。
  • trl — 后训练的另一条路:trl 用 RLHF/DPO 这类偏好优化调生成模型,本库用对比学习调表示模型。两者都基于 HF Trainer 子类化,但 trl 的 loss 吃 logits、本库的 loss 吃嵌入向量——BaseTrainer.compute_loss(base/trainer.py:467)和 trl 的对应物对照着读很有意思。
  • haystack / llamaindex — 下游 RAG 框架:它们把「嵌入模型」当作可插拔组件消费(本库是最常见的后端之一),关心的是 pipeline/索引/检索编排;本库关心的是这个组件本身怎么造、怎么训。

7. 代码地图(导航索引)

主题文件路径符号名
四族共享基类(加载/保存/前向循环)sentence_transformers/base/model.pyBaseModelforward_load_modulessave
稠密双塔主类sentence_transformers/sentence_transformer/model.pySentenceTransformerencodesimilarity
AutoModel 封装(预处理+前向)sentence_transformers/base/modules/transformer.pyTransformerpreprocessforward
池化(mean/cls/max/lasttoken,含 FA2 展平路径)sentence_transformers/sentence_transformer/modules/pooling.pyPooling_forward_padded_forward_flattened
路由模块(query/document、多模态)sentence_transformers/base/modules/router.pyRouterfor_query_document
训练器骨架sentence_transformers/base/trainer.pyBaseTrainercompute_losscollect_featuresget_batch_sampler
数据整理器(列→features)sentence_transformers/base/data_collator.pyBaseDataCollatormaybe_warn_about_column_order
in-batch 负例对比损失(InfoNCE)sentence_transformers/sentence_transformer/losses/multiple_negatives_ranking.pyMultipleNegativesRankingLosscompute_loss_from_embeddings
默认损失(带分数标签的排序)sentence_transformers/sentence_transformer/losses/cosent.pyCoSENTLoss
大 batch 梯度缓存sentence_transformers/base/losses/gradcache.pyCachedLossMixinforward_cachedRandContext
列合并前向sentence_transformers/base/losses/merged_forward.pymerge_feature_batchesembed_columns
batch 采样策略sentence_transformers/base/sampler.pyBatchSamplersNoDuplicatesBatchSamplerGroupByLabelBatchSampler
难负例挖掘sentence_transformers/util/hard_negatives.pymine_hard_negatives
逐对打分模型sentence_transformers/cross_encoder/model.pyCrossEncoderpredictrank
稀疏(SPLADE)sentence_transformers/sparse_encoder/model.py · sparse_encoder/modules/splade_pooling.py · sparse_encoder/losses/splade.pySparseEncoderSpladePoolingSpladeLossFlopsLoss
多向量(ColBERT)sentence_transformers/multi_vector_encoder/model.py · multi_vector_encoder/scoring/colbert.pyMultiVectorEncodercolbert_scoresMultiVectorMask
相似度函数族(含 MaxSim)sentence_transformers/util/similarity.pySimilarityFunctionmaxsimcos_sim
语义搜索工具sentence_transformers/util/retrieval.pysemantic_searchparaphrase_mining
训练示例(MS MARCO 全家桶)examples/sentence_transformer/training/ms_marco/train_bi_encoder_mnrl.py