跳到主要内容

数据截至 (上游 commit 5cc889d47547)

03 · 数据集契约与采样策略

这一章讲什么: 喂给 trainer 的数据长什么样、batch 怎么组(这直接决定对比损失的质量)、以及怎么用 mine_hard_negatives 把数据集从「只有正例」升级成「带难负例」。


1. 它要解决的小问题

对比损失的信号几乎全在「batch 里装了谁」里:

  • batch 里混进重复文本 → 同一个句子既是正例又是别人的负例,目标自相矛盾。
  • 做 batch 内三元组挖掘(BatchHardTripletLoss 等)→ batch 里必须同时有多个类别、每类至少 2 条,否则挖不出三元组。
  • 多个数据集混训 → 各数据集 batch 以什么比例、什么顺序出现。

这些都不是模型问题,是采样问题。本章的三个部件分别回答:数据集契约(列怎么摆)、batch sampler 族(batch 怎么组)、难负例挖掘(负例从哪来)。


2. 数据集契约:列的顺序即语义

训练数据就是一个 datasets.Dataset,每个文本列依次成为损失的第 1、2、3… 个输入。契约有两条:

  1. 列序决定角色。 MNRL 永远把第 1 列当锚、第 2 列当正例、其余列当显式负例——不看列名。docstring 里的原话例子:列是 ["answer", "question"] 时,模型优化的是「给答案猜问题」(sentence_transformers/base/data_collator.py:25-28)。
  2. 标签列按名字认。 label/labels/score/scores 之一会被 collator 抽成 labels 张量(data_collator.py:14DEFAULT_LABEL_COLUMNS:114-118),不参与编码。

护栏有两道,都是「把沉默的事故变成日志」:

  • 列序警告:maybe_warn_about_column_order(data_collator.py:150-187)内置一张常见列名→期望位置的表(anchor→0positive→1negative→2question→0…),发现错位就警告并给出 select_columns 修复建议。
  • ID 列警告:列名是 id 或以 _id/_ids/_idx 结尾时警告「这列会被当文本训练」,并指向 resolve_ids 工具把 ID 解析回文本(data_collator.py:123-129)。

多数据集(DatasetDict 或 dict)混训时,trainer 会自动给样本加 dataset_name 列(base/trainer.py:1135add_dataset_name_column),下游用它选 per-dataset 的 loss、prompt 与 router 映射。


3. 四种 batch sampler

通过 SentenceTransformerTrainingArguments(batch_sampler=...) 选择,分派逻辑在 BaseTrainer.get_batch_sampler(sentence_transformers/base/trainer.py:716-779)。

sampler枚举值解决什么实现
DefaultBatchSamplerBATCH_SAMPLER(默认)无,就是 PyTorch 语义base/sampler.py:195
NoDuplicatesBatchSamplerNO_DUPLICATESbatch 内任何两样本的任何列不撞文本base/sampler.py:406
同上 + 预哈希NO_DUPLICATES_HASHED同上,但先离线把全库哈希成 int64 矩阵base/sampler.py:485
GroupByLabelBatchSamplerGROUP_BY_LABEL每个 batch 保证 ≥2 个标签、每标签 ≥2 样本base/sampler.py:228

3.1 NoDuplicates:把冲突样本「顺延」而不是「丢弃」

朴素做法(老版 NoDuplicatesDataLoader)是换个起点重抽,抽到够为止——既慢又可能凑不齐。现在的实现(NoDuplicatesBatchSampler.__iter__,base/sampler.py:531-611)更巧:

  • 先做一次全库随机排列,再用一个 next_positions 数组把排列串成单链表(:573-576)。
  • 组 batch 时沿链表走:样本与当前 batch 的值集合有交集(任何列撞上)就跳过但留在链表上,顺延给后面的 batch(:588-592);没冲突就从链表摘下、收进 batch。
  • 注意它检查的是跨列的值集合——锚列文本撞上别人的正例列也算冲突,这正是 in-batch negatives 事故的全集。

哈希版(precompute_hashes=True)先用 datasets.map 多进程把每行所有列哈希成 int64(_build_hashes,base/sampler.py:485-529),之后逐 epoch 复用。对图像/音频数据集特别值:不预哈希的话每次查重都要重新解码媒体(docstring,base/sampler.py:449-451)。

3.2 GroupByLabel:为 batch 内三元组挖掘备料

BatchHardTripletLoss 一族要在 batch 内找「同类最难正例、异类最难负例」,前提是 batch 里类别够多。GroupByLabelBatchSampler(base/sampler.py:228)的做法:

  1. 按标签分组,丢掉样本数 <2 的标签,每组裁成偶数(:282-291);合格标签不足 2 个直接 ValueError(:292-296)。
  2. 组内打乱,然后轮转发放:每轮每个存活标签吐 2 个样本,凑满 batch_size 就出一个 batch(__iter__,:313-341);每轮的标签访问顺序再洗牌一次,保证 batch 组合多样。
  3. 流在「只剩 1 个标签」时停止,所以最大标签的尾部会进余数——流长可预先算出来(:298-301,__len__ 据此实现,:343-347)。

要求 batch_size 是 ≥4 的偶数(:279-280)——「≥2 标签 × 每标签 2 样本」的最小公倍。

3.3 多数据集:比例还是轮转

multi_dataset_batch_sampler 控制各数据集 batch 的出现顺序(base/trainer.py:781-823 分派):

策略行为代价
PROPORTIONAL(默认)按数据集大小比例抽样所有样本都会被用到
ROUND_ROBIN轮流从各数据集取,一个耗尽即停每个数据集被平均覆盖,但大数据集用不完

实现分别在 RoundRobinBatchSampler(base/sampler.py:668)与 ProportionalBatchSampler(base/sampler.py:701)。


4. 难负例挖掘:mine_hard_negatives

4.1 为什么需要它

只有 (query, 正文档) 对时,MNRL 的负例全是「别人的正例」——大多很容易区分。难负例(hard negatives)是「看起来很像锚点、但其实不是」的文本,它们提供的梯度信号最强;NV-Retriever 论文(huggingface.co/papers/2407.15831)系统验证了这一点,本函数的 docstring 直接给了论文最佳配置(sentence_transformers/util/hard_negatives.py:89-105)。

4.2 流程

(anchor, positive) 数据集

▼ ① 用现成 SentenceTransformer 编码锚与候选库
▼ ② 相似度检索,取每个锚的 top 候选(可选 FAISS 加速)
▼ ③ 过滤:排除真阳性(range_min 跳过最像的若干条、
│ max_score / absolute_margin / relative_margin 卡阈值)
▼ ④ 可选:用 CrossEncoder 对候选重打分(cross_encoder= 参数)
▼ ⑤ 采样 num_negatives 条(top 或 random)

带负例的数据集(四种输出格式)

实现入口 mine_hard_negatives(sentence_transformers/util/hard_negatives.py:25)。

4.3 关键参数的语义

参数干什么典型用法
range_min跳过最相似的前 k 条候选防「其实是正例」被误标为负例
relative_margin=0.05负例分数必须比正例低至少 5%NV-Retriever 的 TopK-PercPos(95%) 配方
max_score / min_score按绝对分数卡候选窗口过滤太像/太不像的
num_negatives每锚采几条负例论文建议 ≤10
sampling_strategy"top"(总取最难的)或 "random"默认 "top"
cross_encoder用 CrossEncoder 重打分候选更准但更慢
cache_folder嵌入结果存盘复用换 prompt 会使缓存失效(docstring 明说)

4.4 输出格式即下游损失的选择

output_format 四档直接对应不同损失(docstring,hard_negatives.py:241-251):

输出格式列结构配什么损失
triplet(默认)(anchor, positive, negative)(Cached)MNRL、TripletLoss
n-tuple(anchor, positive, negative_1…n)MNRL 的多负例列形态
labeled-pair(anchor, doc, 0/1)CrossEncoder 的 BinaryCrossEntropyLoss
labeled-list(anchor, [docs], [labels])CrossEncoder 的 LambdaLoss

output_scores=True 时把 0/1 标签换成真实分数——蒸馏类损失(MarginMSELossSparseMarginMSELoss)需要的就是分数而非标签。


5. 关键细节 / 坑

  • 默认 sampler 不开去重。 MNRL 官方推荐 NO_DUPLICATES(docstring,multiple_negatives_ranking.py:125-126),但默认是 BATCH_SAMPLER——要自己开。
  • NO_DUPLICATES 会少出样本。 被顺延的冲突样本可能到 epoch 末还留在链表上;__len__ 也因此只是近似值(base/sampler.py:613-614 附近自述 "approximate")。这是「不丢弃就凑不齐整 batch」与「用尽数据」之间的取舍。
  • IterableDataset 不能用这些 sampler。 分派逻辑直接跳过并警告(base/trainer.py:761-766)。
  • GROUP_BY_LABEL 要求数据有标签列,且至少 2 个标签各有 ≥2 样本,否则初始化即报错(base/sampler.py:292-296)。
  • 挖掘的负例质量取决于挖掘用的模型。 用一个弱模型挖,挖到的「难负例」可能只是噪声;流程上常见的是「现成开源模型挖一轮 → 训出更强模型 → 再挖一轮」。
  • 列顺序事故是本库最高发的训练错误之一,好在 collator 的两道警告(§2)基本都拦得住——训练前看一眼日志。

6. 代码地图(本章)

主题文件路径符号名
列契约与警告sentence_transformers/base/data_collator.pyBaseDataCollatormaybe_warn_about_column_orderDEFAULT_LABEL_COLUMNS
sampler 分派sentence_transformers/base/trainer.pyget_batch_samplerget_multi_dataset_batch_sampler
无重复采样sentence_transformers/base/sampler.pyNoDuplicatesBatchSampler_hash_batch
按标签分组sentence_transformers/base/sampler.pyGroupByLabelBatchSampler
多数据集混合sentence_transformers/base/sampler.pyRoundRobinBatchSamplerProportionalBatchSampler
难负例挖掘sentence_transformers/util/hard_negatives.pymine_hard_negatives
ID 解析工具sentence_transformers/util/dataset.pyresolve_ids
训练参数(采样相关字段)sentence_transformers/base/training_args.pyBatchSamplersMultiDatasetBatchSamplers
MS MARCO 难负例配方examples/sentence_transformer/training/ms_marco/train_bi_encoder_mnrl.py全文