Embedding与Reranker微调#
一句话答案#
当评测显示相关文档在库里、却因为领域词或表述差异排不进 top-K 时,才值得微调检索模型。embedding 用 InfoNCE 对比学习:正例相似度在「正例 + 批内负例 + 难负例」上做温度 softmax,难点在负例——难负例要挖,也要过滤假负例、混入随机负例,同一 query 的多个正例不能落进同一 batch 互当负例。reranker 用 pointwise BCE 或 listwise softmax 训练,有分级标签时把档位变成目标分布;上线前先确认线上用的是它的排序还是分数阈值。
核心要点
对比学习、in-batch negatives、池化和归一化的基本概念见 Embedding原理与选型;cross-encoder 的原理、pointwise/pairwise/listwise 三种损失的定义和 ListNet 代码见 Reranker重排原理与选型;NDCG、Recall@K 的定义见 检索评测指标与离线评测。本篇讲微调时的具体做法和容易出问题的地方。
1. 什么时候值得微调#
先排除别的原因,再决定训不训:
| 现象 | 先查什么 | 微调能不能解决 |
|---|---|---|
| 相关文档根本没进候选 | 分块是否切断了关键信息、有没有 BM25 腿(见 混合检索与RRF融合) | 分块和混合检索先改 |
| 领域词、缩写、商品型号召回差 | 通用模型没见过这些词的语义关系 | 适合微调 embedding |
| 召回到了但排在后面 | 候选池里的顺序 | 适合微调 reranker |
| query 形态和文档差异大(口语 vs 书面) | query 改写是否能补 | 微调或改写都可以 |
| 离线指标量不出差异 | 评测集本身(标注覆盖率、K 是否太小) | 先修评测,不然训了也看不出来 |
最后一行最容易被忽略:库里大量文档没有标注时,K 太小的 Recall 会接近随机波动,训练前后的差异被噪声淹掉。先把尺子调到有区分度,再开始训。
2. 数据:三元组从哪来#
训练样本是 (query, 正例, 若干负例)。来源按可靠程度:
- 人工标注:最可靠,量少。分级标注(如 ESCI 的 E 完全匹配 / S 替代品 / C 配件 / I 无关)比二值标注信息多
- 行为日志:点击、购买、停留。量大但有位置偏差(排在前面的更容易被点),未点击不等于不相关
- 合成 query:让 LLM 为每篇文档生成它能回答的问题,按难度分档;生成和过滤方法见 合成数据与拒绝采样
多正例要全展开:一个 query 往往有多个正例。只取第一个会丢掉大量标注;展开成多条 (q, p_i) 样本后,数据量能成倍增长。但展开后同一 query 的样本在文件里是相邻的,顺序读入就会落进同一 batch,下一节会讲为什么这是问题。
3. InfoNCE:公式里每个量在干什么#
对 query q,正例 p,负例集合 N(批内其他样本的正例 + 显式难负例),s 为归一化向量的内积:
温度 τ:对负例 n 的梯度权重正比于它在 softmax 里的概率。τ 小,分布尖锐,梯度集中在最相似的几个负例上,模型更努力区分难负例;τ 太小时,一两个假负例就能主导梯度,训练不稳。τ 大则各负例权重接近平均,学不到细粒度区分。BGE 系列微调常用 0.02 左右,sentence-transformers 里对应参数是 scale = 1/τ。
in-batch negatives:B 条样本的 query 和 B 个正例两两算相似度得到 B×B 矩阵,对角线是正例,其余都是负例,一次前向拿到 B−1 个负例。batch 越大负例越多,所以对比学习偏好大 batch;显存不够时可以跨卡 gather 负例,或用 GradCache 这类方法分块算梯度。
批内冲突:
- 同一 query 的另一个正例被放在同一 batch → 它会被当成负例,模型被要求把真正相关的文档推远
- 不同 query 共享同一个正例文档 → 同样被互当负例
解法是打散样本顺序,并在 loss 里把这些位置 mask 掉:
import torch
import torch.nn.functional as F
def infonce(q, p, hard, qids, pids, tau: float = 0.02):
"""q/p: [B, d],hard: [B, k, d],均已 L2 归一化;qids/pids: [B] 用来找批内冲突"""
B = q.size(0)
s_batch = q @ p.T # [B, B],对角线是正例
s_hard = torch.einsum("bd,bkd->bk", q, hard) # [B, k]
clash = (qids[:, None] == qids[None, :]) | (pids[:, None] == pids[None, :])
clash &= ~torch.eye(B, dtype=torch.bool, device=q.device)
s_batch = s_batch.masked_fill(clash, float("-inf")) # 冲突位置不当负例
logits = torch.cat([s_batch, s_hard], dim=1) / tau
return F.cross_entropy(logits, torch.arange(B, device=q.device))python4. 难负例:挖、过滤、混合#
怎么挖:用基座模型(或上一轮模型)对每个 query 做 ANN 检索,去掉已标注的正例,从排名靠前的一段里采样当难负例。通常跳过最前面几名——排得最靠前又没标注的,很可能是漏标的正例。
假负例过滤:用更强的 cross-encoder 给挖出的候选打分,剔除疑似正例。两种判据:
| 判据 | 做法 | 问题 |
|---|---|---|
| 绝对阈值 | cross-encoder 分数高于某个值就剔除 | 同一模型在不同档位上的分数可能重叠甚至倒挂,会把人工确认过的 S/C 也删掉 |
| query 内相对 | 候选分数高于该 query 已知正例的最高分才剔除 | 更保守,保留的噪声多一些 |
混合随机负例:只用难负例时,模型没见过和 query 完全无关的样本,跨类别的判别能力反而会变差。难负例、随机负例、标注负例按比例混合。
难负例不是越多越好:难负例占比高,模型对近义干扰的抑制变强,但主召回指标可能下降,这是一个取舍,要在评测集上同时看两类指标(比如「召回了多少正例」和「召回了多少配件类干扰」)再定比例。
5. Reranker 微调:数据组织和分级标签#
候选怎么来:用线上召回模型的真实排名取候选,分层采样(比如前几十名、中段、尾段各取一些),让训练分布贴近 reranker 线上实际看到的候选。每组 = 1 个 query + 若干候选(group size),组内至少有一个正例。
分级标签的三种用法:
| 用法 | 做法 | 适合 |
|---|---|---|
| listwise 目标分布 | 档位映射成增益,组内 softmax 成目标分布,和模型分布做交叉熵(ListNet) | 主要消费排序的场景 |
| pointwise 软标签 | 档位映射成 [0,1] 的目标值做 BCE | 要用分数卡阈值 |
| pairwise 跨档对 | 只在不同档位之间构造「高档应排在低档前」的对 | 标签噪声大、只信相对顺序 |
ApproxNDCG 这类直接近似排序指标的损失,在每组只有一个最高档样本时,DCG 主要由它的位置决定,S 与 C 之间相对顺序的影响被明显压低,档位信息利用不足;候选少时近似排名的梯度也弱。它是否优于 ListNet 要做对照实验,不能直接照搬。
分数校准:纯 listwise 训练只约束组内相对大小,分数整体平移不改变损失,训完后分数的绝对含义会变。线上如果用 reranker 分数卡阈值,要:① 加一项 pointwise BCE(见 Reranker重排原理与选型 的 mixed_loss);② 在各档位的标注样本上重新标定阈值,看同等误杀率下能挡住多少无关样本。
query 形态要一致:训练时 query 是完整的意图句,线上送进 reranker 的却是改写后的粗品类词,模型在线上看到的是没见过的输入分布。训练前先看线上实际传给模型的是什么。
6. 评测与防过拟合#
- 按 query 切分训练/验证/测试集,不能按 (q, p) 对随机切,否则同一 query 的正例会同时出现在两侧
- 全库检索评测:embedding 评测要在完整语料上检索,不能只在标注过的候选池里排序;K 要大到指标有区分度
- 编码口径一致:池化方式(CLS 或 mean)、是否 L2 归一化、max_length、query 指令前缀,评测脚本、语料编码、线上服务三处必须一致。不一致不会报错,只是结果变差
- 遗忘检查:只用一种语言或一个领域训练时,另拿一份其他语言/通用领域的评测集,确认没有明显退化
- 过拟合信号:训练 loss 持续降、验证集 Recall 先升后降。embedding 模型一般几亿参数,全参微调常用较小学习率(1e-5 量级)、1–3 个 epoch,按验证集早停
- 端到端验证:离线涨分后,还要在下游链路上做 A/B,并给出置信区间;样本少时涨跌落在区间内只能叫「没测出差异」
- 上线成本:换 embedding 模型要把全部语料重新编码、重建索引。通常建新 collection,双份并存,切换和回滚只改配置
7. 工具链#
| 工具 | 用途 | 备注 |
|---|---|---|
| sentence-transformers | embedding、cross-encoder(reranker)和稀疏编码器训练 | MultipleNegativesRankingLoss 即带 in-batch negatives 的 InfoNCE;CrossEncoderTrainer 训 reranker;util.mine_hard_negatives 挖难负例。v6 调整了部分内部模块的导入路径(旧路径暂时可用但会告警),API 以官方文档为准 |
| FlagEmbedding | BGE 官方仓库,含难负例挖掘、embedder / reranker 微调脚本 | 超参配方可参考其示例 |
| ms-swift | 统一的训练框架,支持 embedding 与 reranker 任务、InfoNCE 损失 | 通过参数和环境变量配置,以文档为准 |
from datasets import Dataset
from sentence_transformers import (SentenceTransformer, SentenceTransformerTrainer,
SentenceTransformerTrainingArguments, losses)
from sentence_transformers.base.sampler import BatchSamplers # v6 路径;v5 及以前是 sentence_transformers.training_args
model = SentenceTransformer("BAAI/bge-m3")
train = Dataset.from_dict({"anchor": queries, "positive": positives, "negative": hard_negs})
loss = losses.MultipleNegativesRankingLoss(model, scale=50.0) # scale = 1/τ,这里 τ = 0.02
args = SentenceTransformerTrainingArguments(
output_dir="out", num_train_epochs=1, per_device_train_batch_size=32, learning_rate=1e-5,
batch_sampler=BatchSamplers.NO_DUPLICATES, # 同一 batch 内不出现重复文本
)
SentenceTransformerTrainer(model=model, args=args, train_dataset=train, loss=loss).train()python面试回答(2分钟版)
微调检索模型的前提是评测证明问题出在模型上:相关文档在库里,但因为领域词、表述差异召回不上或排不上去;分块和混合检索的问题要先排除,评测集本身也要有区分度,否则训了也量不出来。embedding 用 InfoNCE:正例相似度除以温度,在正例加所有负例上做 softmax。温度越小梯度越集中在最难的负例上,BGE 常用 0.02 左右,太小会被假负例带偏。负例来自两处:批内其他样本的正例,和用基座模型 ANN 挖出来的难负例。难负例要过滤假负例,比较稳妥的是用 query 内的相对判据,分数高于已知正例才剔除;还要混随机负例,否则模型没见过完全无关的样本。另外多正例展开后必须打散,不然同一 query 的正例在同一 batch 里会互当负例,loss 里也要把冲突位置 mask 掉。reranker 从召回的真实排名里分层取候选,一组一正多负;有分级标签时,listwise 把档位变成目标分布,pointwise 用软标签 BCE。纯 listwise 会破坏分数的绝对含义,线上要卡阈值就得加 pointwise 项并重新标定。评测按 query 切分、全库检索、三处编码口径一致,再做端到端 A/B。最后要先看线上怎么消费模型输出:如果只拿 reranker 分数做阈值判别,离线排序指标涨了也不一定该上线。结合项目时可以讲:哪一个改动真正带来了收益、用什么评测证明、为什么有的模型训好了却没上线。
追问与易错
追问方向:
- “batch size 对对比学习有什么影响?” → in-batch 负例数等于 B−1,batch 越大负例越多、梯度估计越好;显存不够可以跨卡 gather 负例,或用 GradCache(sentence-transformers 里的
CachedMultipleNegativesRankingLoss)分块计算。 - “多正例展开后为什么必须打散?” → 同一 query 展开的样本相邻写出,顺序读入必然进同一 batch,它的另一个正例会作为批内负例出现在分母里,模型被要求把相关文档推远。打散顺序,或用
BatchSamplers.NO_DUPLICATES、loss 里按 query id mask。 - “假负例怎么发现?” → 抽样人工看难负例里排名最靠前的那批;用更强的 cross-encoder 打分,分数高于该 query 已知正例的候选视为疑似正例剔除。绝对阈值要先看各档位的分数分布是否重叠。
- “embedding 用 LoRA 还是全参?” → embedding 模型通常只有几亿参数,全参微调显存可以承受,主流配方多是全参加小学习率;LoRA 适合模型大、显存紧的情况,需要对照实验确认效果。
- “微调后怎么上线?” → 用新模型重编码全部语料写进新的向量 collection,线上服务通过配置切换 query 编码模型和 collection,两者必须同时切;保留旧 collection 用于回滚。
- “reranker 离线 NDCG 涨了,为什么还可能不能上线?” → 先查线上怎么消费它的分数:如果只拿分数做阈值判别、不依赖排序,排序变好没有收益,而 listwise 训练可能让判别变差。上线前要在各档位标注样本上重新标定阈值并比较误杀率和拦截率。
- “只用英文数据训练,中文会不会退化?” → 可能,要用中英配对的评测集(同一批 query 的两种语言)分别看指标变化。加入中文合成数据时要注意它有没有分级标签,没有的话会稀释分级信号。
易错点:
- ❌ “难负例越多越好” → 难负例占比过高会压低主召回,还会引入更多假负例,要和随机负例混合并看两类指标
- ❌ “换了 embedding 模型只改 query 侧就行” → 文档向量也必须用新模型重编码,否则 query 和文档不在同一个向量空间
- ❌ “listwise 训出来的 reranker 分数可以直接沿用旧阈值” → 分数的绝对刻度变了,阈值必须重新标定