拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

train-sentence-transformers - losses_sparse_encoder

train-sentence-transformers - losses_sparse_encoder

稀疏编码器损失函数(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是一个好的起始比例。
返回列表