跳到主要内容

数据截至 (上游 commit 5cc889d47547)

01 · 模块链与 encode 主线

这一章讲什么: SentenceTransformer 的地基——为什么一个模型是一条 nn.Sequential 模块链,模块之间用什么协议对话,encode() 一次调用从字符串到 numpy 向量中间到底发生了什么,以及模型怎么把自己「自描述」地存盘、再被重建。


1. 它要解决的小问题

把「BERT 类模型」变成「嵌入模型」,差的不只是训练,还有一组固定的流水线动作:tokenize → 过骨干网络 → 把变长的 token 向量压成定长向量(池化)→ 可能再投影/归一化。

这些步骤每个项目都要写,但写法各异、难以复用。Sentence Transformers 的解法是:把每个动作做成一个可插拔模块,模型就是模块的序列。换个骨干、换种池化、尾部加个投影层,都只是换链条上的一节。


2. 思路:字典进、字典出的模块协议

关键设计不在「链」(nn.Sequential 谁都会用),而在链上流动的东西

模块间传递的不是张量,而是一个 features: dict。每个模块的 forward 读字典里自己要的键、写入新键、原样传下去(sentence_transformers/base/modules/module.py:87,抽象方法 Module.forward)。常见键在 docstring 里列得很清楚:input_ids / attention_mask 进来,token_embeddingssentence_embedding 陆续出现。

直觉一句话:features 字典是流水线上的「工件托盘」——每个工位(模块)从托盘上拿料、把加工结果放回托盘,托盘本身越来越满。

# 示意,非源码
features = {"input_ids": ..., "attention_mask": ...} # preprocess 产出
features = transformer(features) # 托盘上多了 "token_embeddings"
features = pooling(features) # 托盘上多了 "sentence_embedding"
features = normalize(features) # "sentence_embedding" 被原地替换为单位向量
return features # 整条链的产出还是这个字典

BaseModel.forward 的真实实现只有六行核心(sentence_transformers/base/model.py:554-570):遍历子模块、按各模块声明的 forward_kwargs 过滤关键字、逐个调用。

这个协议带来两个直接好处:

  • 新增模块零摩擦:写一个 forward(features) -> featuresnn.Module 就能插进链里,上下游都不用改。
  • 训练时损失函数能中途取货:比如 AdaptiveLayerLoss 要看所有层的输出,Transformerall_layer_embeddings 也放进托盘(base/modules/transformer.py:1694),下游谁需要谁取。

3. 链条上的标准模块

3.1 一节链长什么样

用代码拼一个模型(等价于加载一个未按 ST 格式保存的 HF 模型时库的默认行为,sentence_transformers/sentence_transformer/model.py:1315_load_default_modules):

# 示意,非源码
from sentence_transformers import SentenceTransformer, models

transformer = Transformer("microsoft/mpnet-base") # 第 0 节:骨干
pooling = Pooling(transformer.get_embedding_dimension(), pooling_mode="mean")
normalize = Normalize()
model = SentenceTransformer(modules=[transformer, pooling, normalize])

注意第 0 节必须是「输入模块」(负责 preprocess),其余节只做 forward。默认链的池化有个细节:若骨干是 CausalLM 架构且确实是因果的,默认用 lasttoken 池化,否则用 mean(model.py:1370-1375)——因为因果注意力下只有最后一个 token 见过全句。

3.2 模块清单

模块读什么键写什么键文件
Transformerinput_idstoken_embeddings(可配)sentence_transformers/base/modules/transformer.py:644
Poolingtoken_embeddingsattention_masksentence_embeddingsentence_transformers/sentence_transformer/modules/pooling.py:70
Dense默认 sentence_embedding同键或新键sentence_transformers/base/modules/dense.py:21
Normalize默认 sentence_embedding同键覆盖sentence_transformers/base/modules/normalize.py:14
Router整包转发给子链取决于子链sentence_transformers/base/modules/router.py:29

各取一段真实实现看协议怎么用。

Transformer 的 forward(transformer.py:1594)按模态配置调底层 HF 模型的方法,把输出写进 features[self.module_output_name](transformer.py:1674)——输出键名是可配置的,这就是 CrossEncoder 让它写 scores、多向量模型让它写 token_embeddings 的机关。

Pooling 的 forward(pooling.py:121-151)的核心只有一行:把六种模式算出的向量 torch.cat 起来写进 features["sentence_embedding"]。六种模式里常用的三个:

模式干什么实现位置
mean按 attention_mask 加权的 token 均值(默认)pooling.py:200-211
cls取 CLS/首 tokenpooling.py:185-194
lasttoken取最后一个非 pad token(因果模型用)pooling.py:232-241

Pooling 还有第二条路径:输入若被 flash-attention 展平成 (1, 总token数, D)(cu_seq_lens_q 在 features 里),就直接在扁平张量上用 index_add/scatter_reduce 分段聚合,不再补回 padding(pooling.py:245_forward_flattened)。

Normalize(normalize.py:39-43)三行:F.normalize(x, p=2, dim=-1) 后写回原键。归一化之后余弦相似度退化成点积,这就是「检索时只存单位向量、用内积即可」的来源。

3.3 Router:一条链长出分支

有些模型 query 和 document 要走不同编码器(非对称检索),或文本和图像走不同预处理。Router 是链上的「道岔」:forward(task, modality) 解析出一条子链,把 features 依次转发(router.py:459-505)。

最常用的是工厂方法 Router.for_query_document(query_modules, document_modules)(router.py:430),配合 model.encode_query() / model.encode_document() 使用——后两者本质就是给 encode 传 task="query" / task="document"


4. 图示:一次 forward 的托盘流转

preprocess Transformer Pooling Normalize
──────────── features ──────────────► ─┐ ─┐ ─┐
{input_ids, │ │ │
attention_mask} ──────────────────────►│ │ │
▼ ▼ ▼
托盘内容: input_ids ────────────► input_ids ───────────► input_ids ─────────► …
attention_mask attention_mask attention_mask
+ token_embeddings + token_embeddings
+ sentence_embedding ──► 归一化后
(mean 池化) 同键覆盖

怎么读这张图: 横轴是链条顺序,纵轴是托盘在该工位之后的内容——只增不减(除了 Normalize 这种故意覆盖的)。


5. 原理演示:自己实现一遍迷你协议

下面 20 行演示「模块链 + 字典托盘」的全部思想(用伪骨干网络代替真 Transformer):

# 示意,非源码
import torch
from torch import nn

class ToyBackbone(nn.Module):
def forward(self, features): # 读 input_ids,写 token_embeddings
ids = features["input_ids"]
features["token_embeddings"] = torch.randn(ids.shape[0], ids.shape[1], 16)
return features

class MeanPool(nn.Module):
def forward(self, features): # 读 token_embeddings,写 sentence_embedding
mask = features["attention_mask"].unsqueeze(-1)
summed = (features["token_embeddings"] * mask).sum(dim=1)
features["sentence_embedding"] = summed / mask.sum(dim=1).clamp(min=1)
return features

model = nn.Sequential(ToyBackbone(), MeanPool())
out = model({"input_ids": torch.ones(4, 8, dtype=torch.long),
"attention_mask": torch.ones(4, 8)})
print(out["sentence_embedding"].shape) # (4, 16)

重点看:每个模块的返回类型和入参类型相同,这是它们能任意排序、任意拼接的原因。Sentence Transformers 的真实模块只是在这之上加了配置存取(config_keys,module.py:63)和 forward_kwargs 过滤。


6. encode 主线:从字符串到 numpy 向量

SentenceTransformer.encode(sentence_transformer/model.py:754)在模块链之外还做了五件事,按执行顺序:

  1. 参数校验与 prompt 解析:model_kwargs 白名单校验(model.py:866-875,多余 kwargs 直接报错);prompt/prompt_name/默认 prompt 三选一(base/model.py:344_resolve_prompt)。
  2. 按长度降序排序:np.argsort([-self._input_length(s) ...])(model.py:925),让相邻 batch 长度接近、padding 浪费最小。若输入可被 flash-attention 展平(无 padding),改用「最长、最短、次长、次短…」交错,平滑各 batch 的显存峰值(base/model.py:487_interleave_sorted_indices,model.py:926-927 调用点)。
  3. 逐 batch:preprocess → 上设备 → 前向:features = self.preprocess(...)(model.py:933)→ batch_to_deviceself(features, **kwargs)(model.py:941)。注意走 __call__ 是为了让 model.compile() 生效。
  4. 后处理:可选 truncate_dim 截断(Matryoshka 模型的用法,model.py:945-948)、normalize_embeddings、转 numpy。
  5. 还原原始顺序 + 可选量化:np.argsort(length_sorted_idx) 逆排序(model.py:976);precision="int8"/"binary" 等走 quantize_embeddings(model.py:978-979)。

多进程编码(device=["cuda:0","cuda:1"]pool=...)在 model.py:881-913 分流,量化特意放在合并之后统一做一次——注释写明原因:分块量化会让各 chunk 的 int8 校准区间不一致(model.py:896)。


7. 保存与加载:架构即配置

7.1 存盘时写了什么

BaseModel.save(base/model.py:684)落盘两类文件:

  • config_sentence_transformers.json:模型级配置——model_type(哪个家族)、prompts、库版本号(base/model.py:764_get_model_config)。
  • modules.json:一个列表,每节一项 {idx, name, path, type}type完整的类导入路径(如 sentence_transformers.sentence_transformer.modules.pooling.Pooling),path 指向该节的子目录(0_Transformer1_Pooling…),配置与权重都在子目录里(base/model.py:717-759)。第 0 节若声明 save_in_root(Transformer 声明了,transformer.py:770)直接存根目录,与 HF 模型目录约定兼容。

7.2 加载时的三路分派

_load_modules(base/model.py:1036)按现场情况分流:

情况走哪条路行为
没有 modules.json_load_default_modules当作裸 HF 模型,现场拼默认链(见 §3.1)
有,且 model_type 与当前类一致_load_config_modules(base/model.py:1116)modules.json 逐个 import 类、读子目录配置重建
有,但 model_type 不同_load_converted_modules(base/model.py:1343)跨家族转换,例如把 MultiVectorEncoder 存档当 SentenceTransformer 读

config_sentence_transformers.json 里还存了作者的环境要求,加载时先校验版本再建模块(base/model.py:1158-1159check_version_requirements)——一个模型文件自带「我要求什么」的自检。

7.3 模块的自描述配置

每个模块类用类变量 config_keys 声明「我的哪些构造参数要进 config」(module.py:63),get_config_dict 按它反射取值(module.py:113)。旧配置的键名迁移由 config_key_renames 静默完成(module.py:176-230load_config),所以几年前的老存档仍能加载。构造完成后还有一个 on_model_ready 钩子(module.py:389)——比如 MultiVectorMask 在这时把词表里的 skiplist 词解析成 token id(见第 5 章)。


8. 关键细节 / 坑

  • dtype 以第 0 节为准。 构造时会把后续模块统一 cast 到第一节参数的 dtype(base/model.py:262-266);device_map 在场时只把后续模块挪到骨干所在设备,不动骨干本身(base/model.py:268-273)。
  • Pooling 排除 prompt 会复制 mask。 include_prompt=False(INSTRUCTOR 类模型)时,_exclude_prompt_from_maskclone() 再置零(pooling.py:154-169)——因为梯度缓存类损失会拿同一个 features 字典再跑一遍前向,原地改会把两次前向的注意力改得不一致。这是个真实踩过的坑,注释里写了。
  • prompt 的优先级:prompt= 实参 > prompt_name= > default_prompt_name;同时给 promptprompt_name 会忽略后者并告警(base/model.py:364-368)。
  • modules=model_name_or_path= 同给时后者赢,modules 被忽略并告警(base/model.py:230-234)——链来自存档,不是来自你的参数。

9. 代码地图(本章)

主题文件路径符号名
模块基类与配置存取sentence_transformers/base/modules/module.pyModuleconfig_keysload_configon_model_ready
前向循环sentence_transformers/base/model.pyBaseModel.forwardpreprocess
骨干封装sentence_transformers/base/modules/transformer.pyTransformer.forwardTransformer.preprocess
池化sentence_transformers/sentence_transformer/modules/pooling.pyPooling.forward_forward_padded_forward_flattened
投影/归一化sentence_transformers/base/modules/dense.py · normalize.pyDenseNormalize
路由sentence_transformers/base/modules/router.pyRouter.forwardfor_query_document
encode 主线sentence_transformers/sentence_transformer/model.pySentenceTransformer.encode
长度排序与交错sentence_transformers/base/model.py_input_length_interleave_sorted_indices
存盘sentence_transformers/base/model.pyBaseModel.save_get_model_config
三路加载sentence_transformers/base/model.py · sentence_transformer/model.py_load_modules_load_config_modules_load_default_modules