数据截至 (上游 commit 5cc889d47547)
02 · 双塔对比训练
这一章讲什么: 训练侧的核心——trainer 怎么把损失接到模型上、默认的 MNRL/InfoNCE 损失到底在优化什么、什么时候换别的损失,以及两个工程放大器(Matryoshka 可变维度、GradCache 大 batch)。
1. 它要解决的小问题
预训练骨干(如 MPNet)的输出空间里,语义相近的句子并不相近。训练嵌入模型的任务就是重塑这个 空间:
- 拉近:语义该近的对(同义句、query 与命中文档)。
- 推远:其余一切——关键是「推远谁」决定了训练信号的质量。
「双塔(siamese)」指锚和候选用同一个模型分别编码,再比较两个向量。这与 CrossEncoder「句对拼在一起过一次模型」相对(第 4 章)。
2. 思路:损失函数持有模型
大多数训练库里,loss 吃模型的 logits;这里反过来:loss 在构造时拿到整个模型,forward 里自己调用它。
# 示意,非源码
class MyLoss(nn.Module):
def __init__(self, model):
super().__init__()
self.model = model # 损失函数持有模型
def forward(self, sentence_features, labels):
embeddings = [self.model(f)["sentence_embedding"] for f in sentence_features]
return some_contrastive_objective(embeddings, labels)
这个倒置的好处:损失可以自由决定编码几次、编码谁——比如 CachedGISTEmbedLoss 会额外用一个 guide 模型挑负例,MatryoshkaLoss 会对同一批嵌入截到多个维度各算一次子损失。「编码策略」是损失的一部分,不是 trainer 的一部分。
2.1 真实接线:一个 batch 的旅程
Dataset 行: {anchor: str, positive: str, (label?)}
│
▼ collator 逐列 preprocess,列名作前缀拍平 base/data_collator.py:131-146
{anchor_input_ids, anchor_attention_mask,
positive_input_ids, positive_attention_mask, label?}
│
▼ collect_features 按前缀拆回两个字典 base/trainer.py:584
[features_anchor, features_positive], labels
│
▼ loss(features, labels) base/trainer.py:507
└─ embed_columns: 两列(尽量)合并成一次前向 base/losses/merged_forward.py:171
│
▼
对比目标 → 标量 loss
三个环节各有一个值得知道的实现:
- collator(
BaseDataCollator.__call__,sentence_transformers/base/data_collator.py:97):对每个文本列调model.preprocess,产出的键全部加上列名_前缀;label/score等标签列单独抽出(data_collator.py:114-118)。 - collect_features(
base/trainer.py:584-625):按_input_ids等后缀反推列前缀,把扁平 batch 重新分组为「每列一个 features 字典」——正是 loss 期望的输入形态。 - compute_loss(
base/trainer.py:467-517):多数据集训练时按dataset_name选对应的 loss;若 loss 返回 dict(分量损失),逐项累计日志后求和(track_loss_components,base/trainer.py:519)。
此外 override_model_in_loss(base/trainer.py:424-438)会在训练开始时把 loss 里存的模型替换成 DDP/compile 包装后的模型——不换的话梯度就传不到真正被优化的那份权重上。这是个隐蔽但致命的接线 点。
2.2 列合并前向:三次变一次
锚/正/负三列分别过模型 = 三次小 batch 前向,GPU 吃不饱。embed_columns(sentence_transformers/base/losses/merged_forward.py:171)先让 merge_feature_batches(同文件 :53)检查各列预处理结果是否同构(键集合一致、形状兼容、元数据相等),同构则拼成一个大 batch 一次前向再 reshape 切回。
两个保守设计:
- 默认第一列(锚)不合并(
embed_columns的separate_first=True,merged_forward.py:190-195):query 列短、document 列长时,合并会把短列也 pad 到长列宽度,反而更慢。 - 任何同构检查失败都回退到逐列前向,并
warning_once说明原因(_separate_forwards,merged_forward.py:21-33)——损失与梯度在两条路径上完全一致(最多差 dropout 采样),性能优化绝不影响正确性。
3. 核心机制:MNRL / InfoNCE
3.1 它优化什么
MultipleNegativesRankingLoss(sentence_transformers/sentence_transformer/losses/multiple_negatives_ranking.py:17)是本 库事实上的标准损失。一句话:对每个锚,让它与正例的相似度高于它与 batch 内所有其他候选的相似度——其他候选不需要单独准备,同 batch 别人的正例就是负例(in-batch negatives)。
它等价于 InfoNCE / SimCSE / 「in-batch negatives 交叉熵」(类 docstring 里自己列了这些名字,multiple_negatives_ranking.py:60-61)。
3.2 图示:batch 即负例池
batch = [(q1, d1), (q2, d2), (q3, d3)] ← 三对正例
相似度矩阵 cos(Q, D):
d1 d2 d3
q1 [ 正例 负例 负例 ] ← 对 q1:d1 该最大
q2 [ 负例 正例 负例 ] ← 对 q2:d2 该最大
q3 [ 负例 负例 正例 ]
loss = 每行 -log( exp(s正·scale) / Σexp(s·scale) ) 的均值
一个 batch 的 B 对样本,每个锚看到 B−1 个免费负例。所以 batch size 是 MNRL 最重要的超参——MS MARCO 官方示例直接写注释「In-batch negatives are the dominant signal in MNRL」并用 256(examples/sentence_transformer/training/ms_marco/train_bi_encoder_mnrl.py:70)。
3.3 原理演示
# 示意,非源码
import torch
import torch.nn.functional as F
def mnrl(queries, docs, scale=20.0):
# queries/docs: (B, D),第 i 行互为正例
scores = queries @ docs.T * scale # (B, B),对角线是正例
return F.cross_entropy(scores, torch.arange(len(queries)))
重点看:整个损失就是一个相似度矩阵 + 一次交叉熵,对角线是标签。真实实现只是把这个骨架推广到多负例列、多方向与多卡。
3.4 真实实现的四个推广
compute_loss_from_embeddings(multiple_negatives_ranking.py:237-338)在骨架之上加了四个旋钮:
- 多负例列:数据集是 (anchor, positive, negative_1, …, negative_n) 时,
docs = embeddings[1:]全部拼进候选池(:256的docs_all = torch.cat(docs, dim=0)),正例位置用索引算出来(:258-262)。 - 方向(directions):默认只算
query_to_doc;可加doc_to_query(对称项)、query_to_query/doc_to_doc(同类型互斥,GTE 风格)——每个方向是相似度矩阵的一块,自相似项填-inf屏蔽(:272)。GTE 改进版就是四个方向全开 +partition_mode="joint"(docstring 的文献对照表,:135-165)。 - 归一化方式(partition_mode):
joint把所有方向拼进一个 softmax 分母(:327-328);per_direction每方向单独 softmax 再平均。后者与同类互斥方向组合会直接报错——因为那些方向的候选池里没有正例,loss 无定义(:206-216,报错信息里解释了原因)。 - 难负例加权(hardness):给负例 logit 加
hardness_strength * stop_grad(cos_sim)(:286-320),越像锚的负例在 softmax 分母里权重越大。三种模式按「罚谁」区分:
| hardness_mode | 罚谁 | 出处(docstring) |
|---|---|---|
in_batch_negatives | 只罚 batch 内负例 | Lan et al. 2025,α=9 |
hard_negatives | 只罚显式难负例 | EmbeddingGemma,α=5 |
all_negatives | 两者都罚 | 两者合并 |
另外 scale 就是温度的倒数(scale=20 ↔ temperature=0.05,docstring :66-68),多卡时 gather_across_devices=True 用 all_gather_with_grad 把各卡嵌入聚成全局负例池、只算本卡锚的损失,梯度还能传回各卡(:246-253)。
4. 损失族全景:按数据形态选
28 个损失文件,选型的第一依据是你有什么数据:
| 你的数据 | 推荐损失 | 文件 |
|---|---|---|
| (锚, 正例) 对,无分数 | MultipleNegativesRankingLoss(检索默认) | losses/multiple_negatives_ranking.py:17 |
| (锚, 正例) 对 + 教师 guide 模型 | GISTEmbedLoss(guide 挑负例) | losses/gist_embed.py:16 |
| (句A, 句B, 相似度分数) | CoSENTLoss(无 loss 实参时的默认值) | losses/cosent.py:14 |
| (句A, 句B, 0/1 标签) | ContrastiveLoss / OnlineContrastiveLoss | losses/contrastive.py:23 |
| (锚, 正例, 负例) 三元组 | TripletLoss(SBERT 原始) | losses/triplet.py:23 |
| (文本, 类别标签) | BatchHardTripletLoss 等 batch 挖掘族(配 GROUP_BY_LABEL,见第 3 章) | losses/batch_hard_triplet.py:62 |
| (query, 正文档, 负文档) + 教师分数 | MarginMSELoss(蒸馏教师的分差) | losses/margin_mse.py:14 |
| 教师分数分布 | DistillKLDivLoss | losses/distill_kl_div.py:19 |
| 句子 + 分类标签(要分类头) | SoftmaxLoss | losses/softmax.py:19 |
| 只有句子(无对) | DenoisingAutoEncoderLoss / ContrastiveTensionLoss 类 | losses/denoising_auto_encoder.py:37 |
两个对比看很有味道:
- TripletLoss(
triplet.py:98-103)就是relu(d(a,p) − d(a,n) + margin)的均值——每对三元组独立,不用 batch 内其他样本;信号弱但数据要求明确。 - CoSENTLoss(
cosent.py:99-112)把 batch 内所有「分数更高的对应该更相似」的偏序关系一次性 logsumexp 起来——不需要二值标签, 只要有相对分数。它也是SentenceTransformerTrainer不传 loss 时的默认(sentence_transformer/trainer.py:163-168)。
5. 放大器一:Matryoshka——一次训练,多种维度
小问题: 嵌入维度越小,存储和比对越便宜,但截断一个只按全维度训练的模型会掉点。
思路: Matryoshka(俄罗斯套娃)表示学习——让嵌入的前 k 维本身也是一个好嵌入。做法是对同一个嵌入分别截到 [768, 512, 256, 128, 64] 各算一次子损失再加权求和。
MatryoshkaLoss(losses/matryoshka.py:116)是个损失修饰器:包住任意基础损失。非缓存损失走 ForwardDecorator,把模型 forward 的输出按各维度截断并缓存复用(matryoshka.py:226-235);梯度缓存类损失则改包 calculate_loss(matryoshka.py:218-224)——因为缓存损失的反传发生在 hook 里,那时 ForwardDecorator 已拆,只能在算损失处截。两种接法的区分本身就说明了第 2 节「损失持有模型」设计的灵活性。
训练出来的模型,用户在推理时 encode(..., truncate_dim=256) 或构造时传 truncate_dim 就能用低维(model.py:945-948 的截断点),几乎不掉点。
6. 放大器二:GradCache——显存常数化的大 batch
小问题: MNRL 的质量随 batch 增大而升,但 batch 内各样本经 softmax 分母互相依赖,不能用梯度累积拆小——拆了负例池就变了。
思路(GradCache 三步):
① 无梯度逐 mini-batch 前向,得到嵌入并 detach ← 显存只够一个 mini-batch 的激活
│
▼
② 整批嵌入算 loss,反传一次到嵌入层 ← 缓存「每个嵌入的梯度」
│ (嵌入只是几个 GB 的张量,没有激活)
▼
③ 带梯度重放每个 mini-batch,把缓存梯度
.backward() 进该 mini-batch 的嵌入 ← 激活峰值仍是一个 mini-batch
真实实现在 CachedLossMixin.forward_cached(sentence_transformers/base/losses/gradcache.py:397),三个细节值得抄:
- mini-batch 边界预先算好再冻结(
gradcache.py:404-407):模块可能在前向时原地改 features,第 ③ 步必须重放第 ① 步的切分。 - 随机状态快照:
RandContext(gradcache.py:31)在 ① 的每个 mini-batch 记录 RNG 状态,③ 重放时恢复——保证两次前向的 dropout 掩码一致,缓存的梯度才对得上。 - 嵌入是
detach().requires_grad_()的叶子(gradcache.py:424):② 的反传到叶子为止,rep.grad就是缓存;③ 用cached_grad.backward()注入。
使用侧几乎无感:CachedMultipleNegativesRankingLoss(model, mini_batch_size=32)(losses/cached_multiple_negatives_ranking.py:26)与 MNRL 同参,只是多了 mini_batch_size(控制显存)——MS MARCO 示例里 256 大 batch + 32 mini-batch 的搭配(train_bi_encoder_mnrl.py:70-71)就是这个模式的教科书用法。代价是约 2 次前向 + 1 次按 mini-batch 的反传,更慢。
7. 关键细节 / 坑
- 不传 loss 的默认是 CoSENTLoss,不是 MNRL。
get_default_loss(sentence_transformer/trainer.py:163)默认 CoSENT——它需要分数标签;你的数据若是无标签对,必须显式传 MNRL,否则训练目标就错了。 - 列序即语义。 MNRL 永远把第 1 列当锚、第 2 列当正例,不管列名;collator 只对反常列序发警告(
base/data_collator.py:150-187)。把 (answer, question) 顺序的数据喂进去,训出来的是「给答案猜 问题」。 - in-batch negatives 怕重复。 同一句子若既是锚又是别人的正例,会被当负例推开——配
BatchSamplers.NO_DUPLICATES(第 3 章)。 per_direction与同类互斥方向不兼容,会直接ValueError(multiple_negatives_ranking.py:206-216)。- DDP 下别忘了模型替换。
override_model_in_loss(base/trainer.py:424)存在的原因就是:loss 里存的若不是 wrapped 模型,梯度回不到优化器。自定义 loss 若自己存了模型引用,要意识到这层。 - Cached 损失更慢。 多跑一遍无梯度前向 + 按 mini-batch 的反传;docstring 引 GradCache 论文的数据是约多 20% 计算时间(
cached_multiple_negatives_ranking.py:58-60)。用它换 batch 规模,不是换速度。
8. 代码地图(本章)
| 主题 | 文件路径 | 符号名 |
|---|---|---|
| 训练主接线 | sentence_transformers/base/trainer.py | BaseTrainer.compute_loss、collect_features、override_model_in_loss、track_loss_components |
| 列→features 整理 | sentence_transformers/base/data_collator.py | BaseDataCollator.__call__ |
| 列合并前向 | sentence_transformers/base/losses/merged_forward.py | embed_columns、merge_feature_batches |
| InfoNCE/MNRL | sentence_transformers/sentence_transformer/losses/multiple_negatives_ranking.py | MultipleNegativesRankingLoss.compute_loss_from_embeddings |
| 默认排序损失 | sentence_transformers/sentence_transformer/losses/cosent.py | CoSENTLoss.compute_loss_from_embeddings |
| 经典三元组/对比 | sentence_transformers/sentence_transformer/losses/triplet.py · contrastive.py | TripletLoss、ContrastiveLoss |
| 套娃损失 | sentence_transformers/sentence_transformer/losses/matryoshka.py | MatryoshkaLoss、ForwardDecorator、CachedLossDecorator |
| 梯度缓存 | sentence_transformers/base/losses/gradcache.py | CachedLossMixin.forward_cached、RandContext |
| 大 batch MNRL | sentence_transformers/sentence_transformer/losses/cached_multiple_negatives_ranking.py | CachedMultipleNegativesRankingLoss |
| guide 模型负例 | sentence_transformers/sentence_transformer/losses/gist_embed.py | GISTEmbedLoss |
| 官方训练配方 | examples/sentence_transformer/training/ms_marco/train_bi_encoder_mnrl.py | 全文 |