数据截至 (上游 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… 个输入。契约有两条:
- 列序决定角色。 MNRL 永远把第 1 列当锚、第 2 列当正例、其余列当显式负例——不看列名。docstring 里的原话例子:列是
["answer", "question"]时,模型优化的是「给答案猜问题」(sentence_transformers/base/data_collator.py:25-28)。 - 标签列按名字认。
label/labels/score/scores之一会被 collator 抽成 labels 张量(data_collator.py:14的DEFAULT_LABEL_COLUMNS与:114-118),不参与编码。
护栏有两道,都是「把沉默的事故变成日志」:
- 列序警告:
maybe_warn_about_column_order(data_collator.py:150-187)内置一张常见列名→期望位置的表(anchor→0、positive→1、negative→2、question→0…),发现错位就警告并给出select_columns修复建议。 - ID 列警告:列名是
id或以_id/_ids/_idx结尾时警告「这列会被当文本训练」,并指向resolve_ids工具把 ID 解析回文本(data_collator.py:123-129)。
多数据集(DatasetDict 或 dict)混训时,trainer 会自动给样本加 dataset_name 列(base/trainer.py:1135 的 add_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 | 枚举值 | 解决什么 | 实现 |
|---|---|---|---|
DefaultBatchSampler | BATCH_SAMPLER(默认) | 无,就是 PyTorch 语义 | base/sampler.py:195 |
NoDuplicatesBatchSampler | NO_DUPLICATES | batch 内任何两样本的任何列不撞文本 | base/sampler.py:406 |
| 同上 + 预哈希 | NO_DUPLICATES_HASHED | 同上,但先离线把全库哈希成 int64 矩阵 | base/sampler.py:485 |
GroupByLabelBatchSampler | GROUP_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)的做法:
- 按标签分组,丢掉样本数 <2 的标签,每组裁成偶数(
:282-291);合格标签不足 2 个直接ValueError(:292-296)。 - 组内打乱,然后轮转发放:每轮每个存活标签吐 2 个样本,凑满 batch_size 就出一个 batch(
__iter__,:313-341);每轮的标签访问顺序再洗牌一次,保证 batch 组合多样。 - 流在「只剩 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 标签换成真实分数——蒸馏类损失(MarginMSELoss、SparseMarginMSELoss)需要的就是分数而非标签。