☰
train-sentence-transformers - model_architectures
2026/9/30 2:30:24 网站建设 项目流程

模型架构(SentenceTransformer)

SentenceTransformer类是模块组成的torch.nn.Sequential。常见形状是Transformer+Pooling(+ 可选的Normalize/Dense),但支持四种不同的架构家族,正确的选择取决于任务。

四个架构家族

家族骨干池化用例
编码器(双向)BERT、RoBERTa、DeBERTa、MPNet、ModernBERT、XLM-Rmean(默认)或cls短/中文本,通用默认
解码器(因果 LLM)Qwen、Llama、Mistral、Gemmalasttoken长上下文、可指令微调、更高的质量上限
静态嵌入StaticEmbedding模块不适用仅 CPU、<10MB、极快
多模态 / RouterVLM 骨干或组合编码器视情况而定文本 + 图像 / 音频 / 视频

下面,每个家族给出具体设置。

编码器模型(默认)

历史默认,通常仍是文本嵌入的正确选择。

fromsentence_transformersimportSentenceTransformer model=SentenceTransformer("microsoft/mpnet-base")# 自动构建:Transformer(feature-extraction) -> Pooling(mean)。

当SentenceTransformer("<checkpoint>")以原始 HF 编码器调用时,它会自动包装 transformer 并添加Pooling(..., pooling_mode="mean")。

要自定义池化或添加模块:

fromsentence_transformersimportSentenceTransformerfromsentence_transformers.sentence_transformer.modulesimportNormalize,Pooling,Transformer transformer=Transformer("answerdotai/ModernBERT-base")pooling=Pooling(transformer.get_embedding_dimension(),pooling_mode="cls")# 或 "mean"、"lasttoken"、...model=SentenceTransformer(modules=[transformer,pooling,Normalize()])

池化模式:

  • mean(默认)—— token 嵌入的平均值,按注意力掩码遮蔽。最强的默认。
  • cls——[CLS]token 的嵌入。如果基座经过 CLS 预训练则有效。
  • max—— 跨 token 的逐元素最大值。罕见。
  • mean_sqrt_len_tokens—— 按 √seq_len 缩放的均值。经验上对某些任务有帮助。
  • weightedmean—— token 位置加权均值。作为非最后 token 的替代方案,对解码器基座有用。
  • lasttoken—— 最后 token 的嵌入。因果 LM 基座必需(见下文解码器一节)。

不要在训练中途切换池化。只选一次。

解码器 / 因果 LLM 模型

在长上下文、指令遵循、多语言方面表现出色。内存消耗大——通常用 LoRA 训练而非全量微调。

两条设置路径,取决于模型是否已为嵌入适配:

# 路径 A:已适配的嵌入检查点(自带正确的模块):fromsentence_transformersimportSentenceTransformer model=SentenceTransformer("Qwen/Qwen3-Embedding-0.6B")# 直接可用# 路径 B:原始解码器 LLM,手动构建流水线:fromsentence_transformersimportSentenceTransformerfromsentence_transformers.sentence_transformer.modulesimportNormalize,Pooling,Transformer transformer=Transformer("Qwen/Qwen2.5-0.5B",transformer_task="text-generation",# 关键:因果注意力,非双向processor_kwargs={"padding_side":"left"},# last-token 池化需要左填充)pooling=Pooling(transformer.get_embedding_dimension(),pooling_mode="lasttoken")model=SentenceTransformer(modules=[transformer,pooling,Normalize()])

在原始解码器上跳过transformer_task="text-generation"或pooling_mode="lasttoken"会得到看起来合理、直到你跑基准测试才发现问题的嵌入。

为什么用 last-token 池化:因果注意力意味着只有最后一个 token 看到了完整序列。对因果模型做均值池化,平均的是只见过前缀的嵌入——结果不能代表整个输入。

训练解码器基座时:

  • 学习率:通常1e-4或更高(不是编码器的2e-5)。
  • 对 >1B 参数的基座,LoRA 几乎总是正确选择;参见../scripts/train_sentence_transformer_with_lora_example.py(其 docstring 涵盖何时使用、超参数、7B+ 的 QLoRA 和适配器共享)。

静态嵌入

StaticEmbedding完全跳过 transformer——每个 token 通过查找表映射到预计算向量。无注意力、无上下文化。

何时使用:

  • CPU 推理、无 GPU、浏览器 / 边缘 / 端侧部署。
  • 需要 <10MB 模型大小。
  • 每个嵌入的延迟预算 <1ms。
  • 拥有 >100 万训练对(上下文化被逐 token 优化取代;这需要数据)。

何时不使用:

  • 任务需要上下文理解(多义词、句法、长程依赖)。
  • 你只有 <10 万训练对——模型学不到足够的东西。

设置:

fromsentence_transformersimportSentenceTransformerfromsentence_transformers.sentence_transformer.modulesimportStaticEmbeddingfromtokenizersimportTokenizer tokenizer=Tokenizer.from_pretrained("google-bert/bert-base-uncased")static_embedding=StaticEmbedding(tokenizer,embedding_dim=512)model=SentenceTransformer(modules=[static_embedding])

在大型对比数据集(100 万+ 对)上用MultipleNegativesRankingLoss训练。

热启动 vs 随机初始化—— 当你有>100 万训练样本时,随机初始化胜过StaticEmbedding.from_model2vec(...)或.from_distillation(...)热启动。数据集较小时,热启动有帮助。

# 对于较小数据集(<10 万),热启动:static_embedding=StaticEmbedding.from_model2vec("minishlab/potion-base-8M")# 或:static_embedding=StaticEmbedding.from_distillation("sentence-transformers/all-MiniLM-L6-v2",vocabulary=list(tokenizer.get_vocab().keys()))

可运行的端到端配方(随机初始化 + MNRL + Matryoshka + bf16 + lr=2e-1)参见../scripts/train_sentence_transformer_static_embedding_example.py,基准测试参见静态嵌入博客文章。

通过 VLM 骨干实现多模态

现代视觉-语言模型可以直接加载并产生联合的文本+图像嵌入:

fromsentence_transformersimportSentenceTransformer model=SentenceTransformer("Qwen/Qwen3-VL-Embedding-2B",model_kwargs={"attn_implementation":"flash_attention_2"},# 不要在这里设置 torch_dtype;参见 training_args.mdprocessor_kwargs={"min_pixels":28*28,"max_pixels":600*600},)# 检查该模型支持哪些模态:print(model.modalities)# ['text', 'image', 'video', 'message']

训练数据可以混合文本、PIL 图像、图像路径/URL、音频以及如{"image": <PIL>, "text": "describe this"}的混合模态字典。数据 collator 通过模型的preprocess方法处理预处理。

安装多模态附加包:pip install "sentence-transformers[image]"(或[audio]、[video])。

精度:以 fp32 加载并给 TrainingArguments 传bf16=True(或fp16=True)——autocast 处理推理路径。不要在model_kwargs中设置torch_dtype="bfloat16":它会把 Adam 状态置于 bf16 并静默降低质量(参见training_args.md)。

通过 Router 实现多模态

不使用单个 VLM 骨干,而是为每种模态组合独立的编码器:

fromsentence_transformersimportSentenceTransformerfromsentence_transformers.sentence_transformer.modulesimportDense,Pooling,Router,Transformer# 文本编码器text_encoder=Transformer("sentence-transformers/all-MiniLM-L6-v2")text_pooling=Pooling(text_encoder.get_embedding_dimension(),pooling_mode="mean")# 投影文本以匹配图像编码器的维度text_projection=Dense(text_encoder.get_embedding_dimension(),768)# 图像编码器(SigLIP 直接输出池化嵌入)image_encoder=Transformer("google/siglip2-base-patch16-224")router=Router(sub_modules={"text":[text_encoder,text_pooling,text_projection],"image":[image_encoder],},)model=SentenceTransformer(modules=[router])

警告:基于 Router 的模型在初始化时嵌入空间不对齐——你必须训练来对齐它们。维度不同时使用Dense投影层。基于任务的路由(查询与文档使用不同的编码器)也通过route_mappings受支持;参见Router的 docstring。

陷阱

  • 解码器基座使用均值池化:静默产生垃圾嵌入。始终使用lasttoken。
  • Router 多模态不训练:独立编码器的嵌入空间在初始化时不对齐。在训练对齐空间的损失之前,不要指望有用的跨模态相似度。
  • 少于 10 万对的 StaticEmbedding:模型学不到足够的东西。要么通过from_model2vec/from_distillation热启动,要么使用常规编码器。
  • 消费级 GPU 上的大型 VLM 骨干:组合 LoRA +attn_implementation="flash_attention_2"。仅用 LoRA 时,你还可以额外传torch_dtype="bfloat16"——bf16 基座权重是冻结的,所以上面精度规则中关于 Adam 状态的问题不适用(LoRA 适配器保持 fp32,因此其优化器状态保持 fp32)。不用 LoRA 时,遵循精度规则:保持权重 fp32,依赖bf16=Trueautocast。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询