稀疏编码器损失函数(SPLADE)
所有损失函数都位于sentence_transformers.sparse_encoder.losses中。
本参考针对SPLADE架构(Transformer + SpladePooling)。稀疏编码器包还为 CSR 架构(Transformer + Pooling + SparseAutoEncoder)导出了CSRLoss和CSRReconstructionLoss;这些不在本文范围内——如果你训练 CSR 模型,请参见 sbert.net 文档。
选择损失意味着 (a) 选择一个基础损失(对比、回归、蒸馏),(b) 将其包装在SpladeLoss中以添加 FLOPS 正则化。
顶层决策表
| 你有 | 使用 |
|---|---|
(anchor, positive)或三元组,SPLADE 架构 | SpladeLoss(loss=SparseMultipleNegativesRankingLoss(model), ...) |
| 相同,想要 256+ 的有效批次大小 | CachedSpladeLoss(...) |
带标签的(text1, text2, score)对 | SparseCoSENTLoss或SparseCosineSimilarityLoss |
| 从交叉编码器教师蒸馏 | SparseMarginMSELoss |
| 列表级蒸馏 | SparseDistillKLDivLoss |
| 显式三元组 | SparseTripletLoss |
核心包装器:SpladeLoss
SpladeLoss在另一个稀疏损失之上添加FLOPS 正则化。FLOPS 正则化惩罚非零激活,保持嵌入真正稀疏。
loss=SpladeLoss(model=model,loss=SparseMultipleNegativesRankingLoss(model=model),query_regularizer_weight=5e-5,document_regularizer_weight=3e-5,)query_regularizer_weight:惩罚查询嵌入中非零项的程度。document_regularizer_weight:对文档同样如此。- 典型范围:1e-5 到 1e-4。更高 = 更稀疏的嵌入、更低的召回率;更低 = 更稠密,可能更好的召回率。
- 当损失是
SpladeLoss时,SparseEncoderTrainer会自动注册一个SpladeRegularizerWeightSchedulerCallback。该回调在训练的前约 33% 将权重从 0 爬升到目标值;默认形状是SchedulerType.QUADRATIC(不是线性的)。爬坡长度和形状在回调上配置(SpladeRegularizerWeightSchedulerCallback(loss=..., warmup_ratio=..., scheduler_type=...)),而不是在SpladeLoss上;要覆盖它,请自己实例化回调并通过callbacks=[...]传入。这个爬坡很重要;从第 0 步就开始完全正则化会扼杀学习。
使用CachedSpladeLoss获取 GradCache 变体。
对比损失(无标签)
SparseMultipleNegativesRankingLoss
双编码器 MNRL 的稀疏对应物。批内对比。
inner=SparseMultipleNegativesRankingLoss(model=model)loss=SpladeLoss(model=model,loss=inner,query_regularizer_weight=5e-5,document_regularizer_weight=3e-5)- 对 SPLADE 架构始终包装在
SpladeLoss中。 - 在训练参数上设置
batch_sampler=BatchSamplers.NO_DUPLICATES。
SparseTripletLoss
在显式(anchor, positive, negative)上的经典三元组边距损失。
带标签回归损失
SparseCoSENTLoss
用于(text1, text2, score)的成对排序损失。镜像双编码器CoSENTLoss。
SparseCosineSimilarityLoss
余弦相似度上的 MSE。更简单,通常比 CoSENT 差。
SparseAnglELoss
复数空间中的基于角度损失。CoSENT 的替代方案。
蒸馏损失
SparseMSELoss
嵌入 MSE。学生稀疏嵌入应与教师嵌入匹配。
- 数据:
(text, teacher_embedding)。 - 教师可以是稠密双编码器或另一个稀疏模型。
SparseMarginMSELoss
来自交叉编码器教师的边距 MSE。
- 数据:
(query, positive, negative, score_diff),其中score_diff = teacher_score(query, positive) - teacher_score(query, negative)。 - 从交叉编码器标签训练 SPLADE 的典型配方(ms-marco 蒸馏)。
- 对 SPLADE 包装在
SpladeLoss(model, loss=SparseMarginMSELoss(model), ...)中。
SparseDistillKLDivLoss
列表级 KL 散度蒸馏——学生在候选上的 softmax 分布应与教师匹配。
独立正则化器
FlopsLoss
独立的 FLOPS 正则化器。通常通过SpladeLoss使用,而不是直接使用。
关于正则化器权重调优和稠密输出恢复,参见troubleshooting.md(“SPLADE embeddings are dense”)。MLM 头要求:base_model_selection.md(SPARSE 一节)。激活维稀疏度目标和监控方法:evaluators_sparse_encoder.md(Sparsity tracking)。
陷阱
- 在 SPLADE 模型上
SparseMultipleNegativesRankingLoss不包SpladeLoss:没有 FLOPS 正则化 -> 稠密输出违背 SPLADE 的目的。始终包装。 CachedSpladeLoss+gradient_checkpointing=True:崩溃。二选一。- 从第 0 步就以完整 FLOPS 正则化开始训练:模型到处输出零并卡住。内置调度器会避免这种情况——除非你知道原因,否则不要覆盖它。
query_regularizer_weight==document_regularizer_weight:通常是错的。查询应比文档更稀疏(每个查询的词更少)。由于更高的正则化会产生更多零,请给查询权重更大的值。query_regularizer_weight=5e-5、document_regularizer_weight=3e-5是一个好的起始比例。