训练 CrossEncoder 重排器:数据、难负例、损失与真实评估

训练 CrossEncoder 重排器:数据、难负例、损失与真实评估

原文:Sentence Transformers 文档贡献者,Cross Encoder Training Overview。本文依据 2026-10-05 读取的完整官方页面及其 Markdown 源文翻译整理。该页面持续更新,没有固定发布版本号;文中区分 v4.0 之前的旧训练接口与当前 Trainer 用法。所有代码只做静态审核,没有下载模型、训练、评估或上传。

为什么微调重排器

在 Retrieve and Rerank 检索系统中,第一阶段先检索出一批候选,再由 CrossEncoder 同时读取查询和每个候选,为它们重新打分排序。训练数据与实际任务不一致时,重排甚至可能降低效果,因此微调往往很重要。重排器每个输入对输出一个分数,即 num_labels=1。

CrossEncoder 也能做文本对分类,例如把自然语言推断的两个句子分为“矛盾、蕴含、中立”。这类任务通常需要多个输出,不能直接套用单分数重排器的损失与标签设置。本文主线是训练第二阶段重排器;它无法找回第一阶段根本没有召回的相关文档。

CrossEncoder 训练与评估流程:问答正例和难负例进入训练器,配合单输出模型、损失及训练参数;独立评估保留检索器实际候选,对比初始排序与重排后的指标。
未完纪原创示意:训练和评估分别准备数据,重排只作用于召回候选。不是训练结果截图。

整个训练涉及四个核心部分:模型、数据集、损失函数和 Trainer;训练参数与评估器虽可省略,通常也应该明确配置。原页面最顶部还出现一条面向 AI coding agents 的技能安装提示。它与 CrossEncoder 训练正文无关,本稿没有转录该可变安装命令;如需采用,应先独立审查安装来源、权限与影响。本文没有安装或执行该技能。

选择模型架构和输出头

CrossEncoder 封装 Transformers 预训练模型,用已有或新加的小型输出头产生输入对分数。继续微调现成 CrossEncoder 通常只需加载它:

from sentence_transformers import CrossEncoder
model = CrossEncoder("cross-encoder/ms-marco-MiniLM-L6-v2")

从 BERT、RoBERTa、ModernBERT 等 encoder / sequence-classification checkpoint 起步时,如果已有分类头则复用,否则自动添加。类别数取模型配置或显式 num_labels:

reranker = CrossEncoder(
    "google-bert/bert-base-uncased",
    num_labels=1,
    model_kwargs={"torch_dtype": "float32"},
)
classifier = CrossEncoder(
    "FacebookAI/xlm-roberta-base",
    num_labels=3,
    model_kwargs={"torch_dtype": "float32"},
)

单输出模型适合重排或回归式任务,例如配合 BinaryCrossEntropyLoss;多类别任务要使用匹配类别数与目标的分类损失。内存允许时,原文建议以 fp32 加载待训练权重,即使随后使用混合精度;直接以低精度存储权重,训练后期小更新可能被舍入。

基础模型应与语言和领域相符。可从 fill-mask、sentence similarity、feature-extraction 模型中寻找候选;原文以韩语任务的 klue/bert-base 对比英文 BERT 说明这一原则,这不是对任意新数据的性能保证。

CrossEncoder 还可包装配置为 ForCausalLM 的生成模型,例如 CrossEncoder("google/gemma-2-2b-it")。它通过 LogitScore 读取“yes/no”或“1/0”等对应 token 的 logits,形成单个分数。此结构用于打分与重排,不能把它视为每类一个 logit 的多类别分类器;多类别任务应优先用 encoder / sequence-classification 架构。模型访问条件、许可、内存与推理成本也需分别核对。

数据集格式比列名更重要

Trainer 接受一个 datasets.Dataset,或多个数据集组成的字典 / DatasetDict。Hub 数据集可直接加载;部分数据集需要指定子集,例如 AllNLI 的 pair、pair-class、pair-score、triplet 对应不同格式:

from datasets import Dataset, load_dataset

train_dataset = load_dataset("sentence-transformers/all-nli", "pair-class", split="train")
eval_dataset = load_dataset("sentence-transformers/all-nli", "pair-class", split="dev")
local_csv = load_dataset("csv", data_files="my_file.csv")
local_json = load_dataset("json", data_files="my_file.json")

# 先完成本地清洗,再把同长度的列交给 Dataset。
anchors = ["示例问题"]
positives = ["相应答案"]
prepared = Dataset.from_dict({"anchor": anchors, "positive": positives})

最后两列只是格式示例,一条样本不足以训练。数据来源可以是 CSV、JSON、Parquet、Arrow 或 SQL 等,具体加载方式按 Datasets 的支持接口选择。

为数据、模型和损失匹配时,检查三件事:

  1. 除了名为 label、labels、score、scores 的列,其余列都会被视为输入。输入数量要符合损失要求。
  2. 需要标签的损失必须存在上述名称之一的标签列。
  3. 模型输出维度要符合损失要求。

非标签列的名字不决定角色,列顺序才决定。如果列为 good_answer、bad_answer、question,三元组损失仍会把第一列当 anchor、第二列当 positive、第三列当 negative,语义就错了。用 select_columns 显式重排,并移除 sample_id、metadata、source、type 等额外列,避免它们被当模型输入。

dataset = dataset.select_columns(["question", "good_answer", "bad_answer"])

例如两列文本加 0–1 标签、单输出模型,可以配 BinaryCrossEntropyLoss。不能仅因为都是“文本对”就把三类别 NLI 标签塞进同样的二元损失。

多模态数据

使用多模态骨干时,输入列也可包含 PIL 图像、文件路径或 URL、音频、视频及多模态字典。非标签列作为输入、label 作为目标的规则仍适用。原文的图文相关性例子为:

from datasets import Dataset
from PIL import Image

dataset = Dataset.from_dict({
    "image": [Image.open("cat.jpg"), Image.open("cat.jpg"), Image.open("dog.jpg")],
    "text": ["a photo of a cat", "a photo of a dog", "a photo of a dog"],
    "label": [1, 0, 1],
})

若要双向学习,原文分别把图像作为输入、文本作候选,以及把文本作为输入、图像作候选;两个子数据集列顺序相反,再以名称映射传给同一个 Trainer:

train_image_to_text = full_dataset.select_columns(["image", "text", "label"])
train_text_to_image = full_dataset.select_columns(["text", "image", "label"])

trainer = CrossEncoderTrainer(
    model=model,
    args=args,
    train_dataset={
        "image_to_text": train_image_to_text,
        "text_to_image": train_text_to_image,
    },
    loss=loss,
)

代码边界:这段是原文 Trainer 主体,不是独立脚本;full_dataset、兼容多模态的 model、训练参数 args 和损失 loss 必须先定义。数据整理器调用模型的 preprocess 处理多模态输入;图像支持需额外安装 sentence-transformers[image]。具体 checkpoint 是否支持图片/文本输入应先核实。

图像路径需由使用者提供,本文没有假造这些图片。图像支持需要相应依赖(原文举 sentence-transformers[image])。数据整理器会通过模型的 preprocess 进行预处理;encoder 与 causal-LM 的数据格式相同,但架构能力仍需匹配。若要双向训练,可将 image、text、label 与 text、image、label 两种列序列分别作为子数据集传入同一 Trainer。

构造软负例与难负例

负例质量经常决定重排效果。软负例与问题几乎无关;难负例看似有关,实际上不能回答。比如问“Apple 在哪里成立”,关于一座桥的文本是软负例,关于富士苹果品种的文本则可能是难负例。

mine_hard_negatives() 可用一个嵌入模型,从问答正例中挖掘候选。原文以 GooAQ 的前十万条问答和 CPU 上的 static-retrieval 模型示范:

from datasets import load_dataset
from sentence_transformers import SentenceTransformer
from sentence_transformers.util import mine_hard_negatives

data = load_dataset("sentence-transformers/gooaq", split="train").select(range(100_000))
embedding_model = SentenceTransformer(
    "sentence-transformers/static-retrieval-mrl-en-v1", device="cpu"
)
hard_data = mine_hard_negatives(
    data,
    embedding_model,
    num_negatives=5,
    range_min=10,
    range_max=100,
    max_score=0.8,
    absolute_margin=0.1,
    relative_margin=0.1,
    sampling_strategy="top",
    batch_size=4096,
    output_format="labeled-pair",
    use_faiss=True,
)

这里跳过最相似的前一段候选,仅在指定排名范围内挑选;max_score 排除过于相似的候选,absolute_margin 要求负例相似度比正例低至少固定量,relative_margin 则要求至少低于正例分数绝对值的一定比例。当前 API 明确相对界限为 positive_score - abs(positive_score) * relative_margin。这些阈值只是示例,不保证选出的每条都是真负例,仍应抽查标注。

labeled-pair 输出 query、passage、label,正例 1、负例 0,适合二元交叉熵。FAISS 需要兼容环境中的额外安装,4096 是原例批大小,不是任意机器都适用。原文示例产生 100,000 个正例与 436,925 个负例,共 536,925 行;未必每个问题都有足额负例。这些是上游示例输出,本文没有复现,不能把其耗时或数量作为保证。

查看上游脚本的完整关键输出(未在本文复现)
Dataset({
    features: ['question', 'answer'],
    num_rows: 100000
})

Batches: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████| 22/22 [00:01<00:00, 12.74it/s]
Batches: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████| 25/25 [00:00<00:00, 37.50it/s]
Querying FAISS index: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████| 7/7 [00:18<00:00,  2.66s/it]
Metric       Positive       Negative     Difference
Count         100,000        436,925
Mean           0.5882         0.4040         0.2157
Median         0.5989         0.4024         0.1836
Std            0.1425         0.0905         0.1013
Min           -0.0514         0.1405         0.1014
25%            0.4993         0.3377         0.1352
50%            0.5989         0.4024         0.1836
75%            0.6888         0.4681         0.2699
Max            0.9748         0.7486         0.7545
Skipped 2,420,871 potential negatives (23.97%) due to the absolute_margin of 0.1.
Skipped 43 potential negatives (0.00%) due to the max_score of 0.8.
Could not find enough negatives for 63075 samples (12.62%). Consider adjusting the range_max, range_min, absolute_margin, relative_margin and max_score parameters if you'd like to find more valid negatives.
Dataset({
    features: ['question', 'answer', 'label'],
    num_rows: 536925
})

{
    'question': 'how to transfer bookmarks from one laptop to another?',
    'answer': 'Using an External Drive Just about any external drive, including a USB thumb drive, or an SD card can be used to transfer your files from one laptop to another. Connect the drive to your old laptop; drag your files to the drive, then disconnect it and transfer the drive contents onto your new laptop.',
    'label': 0
}

Difference 列的计算口径:它是对保留下来的正例—负例配对逐对计算差值,再对所有配对取均值。不同问题保留的负例数量可能不同,正例分数因而按各自负例数重复参与;所以该均值不一定等于总体 Mean(Positive) - Mean(Negative)。这是上游实现的静态读数,本文未重新运行挖掘;可对照当前实现。

选择损失与训练参数

损失把一个批次的错误量化为优化目标。没有对所有任务最好的单一损失,选择取决于现有数据与目标。MultipleNegativesRankingLoss 可以接收相关文本对或三元组;CachedMultipleNegativesRankingLoss 通过小批处理控制内存。BinaryCrossEntropyLoss 接收带标签文本对,虽然比 LambdaLoss、ListNetLoss 等学习排序方法简单,仍是很有竞争力的选择。

from sentence_transformers import CrossEncoder
from sentence_transformers.cross_encoder.losses import MultipleNegativesRankingLoss

model = CrossEncoder(
    "FacebookAI/xlm-roberta-base", num_labels=1,
    model_kwargs={"torch_dtype": "float32"},
)
loss = MultipleNegativesRankingLoss(model)

训练参数一组控制优化和资源:learning_rate、lr_scheduler_type、warmup_steps、num_train_epochs / max_steps、每设备批大小、自动批大小、fp16 / bf16、梯度累积、梯度检查点、优化器、数据加载 worker 及预取等。另一组控制观察与保存:eval_strategy / eval_steps、save_strategy / save_steps / save_total_limit、logging_steps、report_to、run_name,以及 Hub 的上传目标和策略。

原文用一个 epoch、学习率 2e-5、warmup_steps=0.1,评估与保存每 0.1 训练步比例、日志每 0.01 比例作示例。比例参数的解释与支持范围应按安装的 TrainingArguments 核对。混合精度需要硬件支持;不能因为示例写 bf16=True 就默认可用。批内负例损失可受益于 BatchSamplers.NO_DUPLICATES;原文局部参数例使用该名称却没给出导入,实际代码要从 sentence_transformers.base.sampler 导入 BatchSamplers。

独立参数示例(补全导入并默认关闭外发):原文局部片段使用了 BatchSamplers.NO_DUPLICATES,却没有导入该名称;以下显式导入当前文档路径,并把混合精度、远程报告与 Hub 写入设为关闭。比例型步骤参数的含义及支持范围仍须按锁定版本检查。此代码未运行。

from sentence_transformers.cross_encoder import CrossEncoderTrainingArguments
from sentence_transformers.base.sampler import BatchSamplers

args = CrossEncoderTrainingArguments(
    output_dir="models/reranker-MiniLM-msmarco-v1",
    num_train_epochs=1,
    per_device_train_batch_size=16,
    per_device_eval_batch_size=16,
    learning_rate=2e-5,
    warmup_steps=0.1,
    fp16=False,  # 仅在硬件支持并已核实后启用混合精度
    bf16=False,
    batch_sampler=BatchSamplers.NO_DUPLICATES,
    eval_strategy="steps",
    eval_steps=0.1,
    save_strategy="steps",
    save_steps=0.1,
    save_total_limit=2,
    logging_steps=0.01,
    report_to="none",  # 关闭默认远程指标接收端
    push_to_hub=False,  # 仅写本地
    run_name="reranker-MiniLM-msmarco-v1",
)

评估器:衡量真正关心的排序

提供 eval_dataset 可在训练中得到评估损失,但排序任务通常还需要 MAP、MRR、NDCG 等更直接指标。评估器可以在训练前、期间和之后运行;可与 eval_dataset 同时使用,也可只选其中一种。执行频率由 eval_strategy 与 eval_steps 决定。

评估器 输入
CrossEncoderClassificationEvaluator 文本对及二元 / 多类别标签。
CrossEncoderCorrelationEvaluator 文本对及相似度分数。
CrossEncoderNanoBEIREvaluator 自动从 Hugging Face 加载相应基准,不要求调用者先提供数据。
CrossEncoderRerankingEvaluator query、positive、negative 或有序 documents 构成的样本。
SequentialEvaluator 把多个评估器组合成一个,交给 Trainer。

NanoBEIR 可直接初始化,或选 msmarco、nfcorpus、nq 等子集;它仍会加载数据并消耗计算,不能把“无必填参数”理解成不依赖外部数据。分布式训练时,普通训练 / 评估数据会分配到设备,但 evaluator 只在第一台设备运行。

不要让评估器补回检索器漏掉的正例

为任务内重排评估挖掘候选时,用 include_positives=True 让正例也参与嵌入检索的排序候选,并用 output_format="n-tuple" 保存顺序。把这些候选作为 documents 交给评估器,就能同时比较原始检索顺序和重排顺序。

关键在于 CrossEncoderRerankingEvaluator 默认可能把全部标注正例补入待重排集合,即使初始召回根本没有找到它们。这样可增强“模型会不会排序”的训练信号,却会高估完整两阶段检索表现。要衡量实际召回链路,应设置:

reranking_evaluator = CrossEncoderRerankingEvaluator(
    samples=samples,
    name="gooaq-dev",
    always_rerank_positives=False,
)

原文在 1,000 个查询上的示例对比为:补入正例时,NDCG@10 从 59.12 到 71.35;只重排实际检索到的正例时,从 59.12 到 70.10。MAP 分别从 53.28 到 67.28 / 66.12,MRR@10 分别从 52.40 到 66.65 / 65.61。这些是原文的具体案例,不是本文测试结果,也不是微调可保证的收益。

完整基准与 GooAQ 任务内重排评估示例

上游页另给出一个直接加载 NanoBEIR 的基准例,以及一个从 GooAQ 留出集挖出检索候选、先看嵌入模型基线再看 CrossEncoder 重排的完整示例。调用会联网获取公开模型、数据和基准;下列代码仅供静态阅读,未在本稿运行。为评估真实两阶段链路,参数显式设为 always_rerank_positives=False;上游默认行为会补入全部标注正例,适合比较模型排序能力,但指标可能高于真实召回上限。

展开 NanoBEIR 最小评估代码
from sentence_transformers import CrossEncoder
from sentence_transformers.cross_encoder.evaluation import CrossEncoderNanoBEIREvaluator

model = CrossEncoder("cross-encoder/ms-marco-MiniLM-L6-v2")
dev_evaluator = CrossEncoderNanoBEIREvaluator()
# 运行时会从 Hugging Face 加载相应基准;本文未执行:
# results = dev_evaluator(model)
展开 GooAQ 候选挖掘与重排完整代码
from datasets import load_dataset
from sentence_transformers import SentenceTransformer
from sentence_transformers.cross_encoder import CrossEncoder
from sentence_transformers.cross_encoder.evaluation import CrossEncoderRerankingEvaluator
from sentence_transformers.util import mine_hard_negatives

model = CrossEncoder("cross-encoder/ms-marco-MiniLM-L6-v2")
full_dataset = load_dataset("sentence-transformers/gooaq", split="train").select(range(100_000))
dataset_dict = full_dataset.train_test_split(test_size=1_000, seed=12)
eval_dataset = dataset_dict["test"]
print(eval_dataset)
# 上游文档显示:Dataset(features=['question', 'answer'], num_rows=1000)

embedding_model = SentenceTransformer(
    "sentence-transformers/static-retrieval-mrl-en-v1", device="cpu"
)
hard_eval_dataset = mine_hard_negatives(
    eval_dataset,
    embedding_model,
    corpus=full_dataset["answer"],
    num_negatives=50,
    batch_size=4096,
    output_format="n-tuple",
    include_positives=True,
    use_faiss=True,
)
print(hard_eval_dataset)
# 1,000 rows; columns are question, answer, and negative_1 through negative_50.

reranking_evaluator = CrossEncoderRerankingEvaluator(
    samples=[
        {
            "query": sample["question"],
            "positive": [sample["answer"]],
            "documents": [
                sample[column_name]
                for column_name in hard_eval_dataset.column_names[2:]
            ],
        }
        for sample in hard_eval_dataset
    ],
    batch_size=32,
    name="gooaq-dev",
    always_rerank_positives=False,
)
# 实际执行会下载数据并计算指标;本文未执行:
# results = reranking_evaluator(model)
查看上游记录的示例评估输出(非本文实测)
CrossEncoderRerankingEvaluator: Evaluating the model on the gooaq-dev dataset:
Queries:  1000     Positives: Min 1.0, Mean 1.0, Max 1.0   Negatives: Min 49.0, Mean 49.1, Max 50.0
          Base  -> Reranked
MAP:      53.28 -> 66.12
MRR@10:   52.40 -> 65.61
NDCG@10:  59.12 -> 70.10

来源页默认补入未召回正例时的同一示例为 MAP 53.28 → 67.28、MRR@10 52.40 → 66.65、NDCG@10 59.12 → 71.35;设置 always_rerank_positives=False 后,上方数值为 66.12、65.61、70.10。两组都是上游文档记载的示例指标,不能当作本文实验结果、版本保证或普遍收益。

相似度与分类任务的其他例子

原文的 STSb 例加载 cross-encoder/stsb-TinyBERT-L4,将 validation 的 sentence1 / sentence2 配对,把 score 传给 CrossEncoderCorrelationEvaluator。AllNLI 分类例则加载 cross-encoder/nli-deberta-v3-base,取 pair-class 子集的前 1,000 个 dev 样本,用 {0: 1, 1: 2, 2: 0} 把数据标签映射到该模型的类别顺序。该映射是模型特定的,换 checkpoint 后不能照抄。

静态修正:原文分类评估块导入了 TripletEvaluator / SimilarityFunction,却调用未导入的 CrossEncoderClassificationEvaluator。应使用:

from sentence_transformers.cross_encoder.evaluation import (
    CrossEncoderClassificationEvaluator,
)

这修复的是可见导入遗漏,未声称其余模型、数据或版本兼容性已经运行验证。

独立评估:STSb 相关性与 AllNLI 分类

原文把相似度回归与自然语言推断分类作为两种独立评估示例。STSb 代码从 validation 读取句对和连续分数,交给相关性评估器;AllNLI 代码从 pair-class 的 dev 前 1,000 条取前提/假设,并按训练该模型时采用的类别顺序映射标签。它们不是重排基准:STSb 看分数相关性,AllNLI 看分类表现,不能把这些指标替代 MAP、MRR 或 NDCG。

应将独立评估数据与训练数据隔离;换模型时必须重新核实标签映射。这里保留原文示例的 checkpoint 与数据切分,不把它们视为新模型的保证或本稿实测结果。

from datasets import load_dataset
from sentence_transformers import CrossEncoder
from sentence_transformers.cross_encoder.evaluation import CrossEncoderCorrelationEvaluator

# 在 STSb validation 上评估句对相似度与人工分数的相关性。
model = CrossEncoder("cross-encoder/stsb-TinyBERT-L4")
eval_dataset = load_dataset("sentence-transformers/stsb", split="validation")
pairs = list(zip(eval_dataset["sentence1"], eval_dataset["sentence2"]))
dev_evaluator = CrossEncoderCorrelationEvaluator(
    sentence_pairs=pairs,
    scores=eval_dataset["score"],
    name="sts_dev",
)
# 单独运行评估(本稿未执行):results = dev_evaluator(model)
from datasets import load_dataset
from sentence_transformers import CrossEncoder
from sentence_transformers.cross_encoder.evaluation import CrossEncoderClassificationEvaluator

# 用 AllNLI 的 dev 子集检查 NLI 分类;类别映射只适用于此 checkpoint。
model = CrossEncoder("cross-encoder/nli-deberta-v3-base")
max_samples = 1_000
eval_dataset = load_dataset(
    "sentence-transformers/all-nli",
    "pair-class",
    split=f"dev[:{max_samples}]",
)
pairs = list(zip(eval_dataset["premise"], eval_dataset["hypothesis"]))
label_mapping = {0: 1, 1: 2, 2: 0}
labels = [label_mapping[label] for label in eval_dataset["label"]]
cls_evaluator = CrossEncoderClassificationEvaluator(
    sentence_pairs=pairs,
    labels=labels,
    name="all-nli-dev",
)
# 单独运行评估(本稿未执行):results = cls_evaluator(model)

把组件接成可审查的训练流程

原文 Simple Example:GooAQ 随机负例完整训练脚本

以下保留原文简版示例的完整本地训练流程:加载 MiniLM、从 GooAQ 固定抽取 100,000 条并拆分 1,000 条留出集、采样随机负例、训练前后跑 NanoBEIR、保存模型。原例末尾把 push_to_hub 放在默认执行路径里;本稿把该远程写入块改成注释,并将旧登录提示更新为当前 CLI 形式。这个差异是安全编辑,不代表本轮运行过脚本。

版本与资源限制:代码来自未固定版本号的持续更新文档;上游示例设 bf16=True,本文将其改为默认关闭,并将 report_to 设为 none,避免未确认时启动 W&B 等远程指标上报;核实硬件、依赖及遥测目的地后再启用。首次运行会下载模型、GooAQ 与 NanoBEIR 基准,需要联网和足够存储;本文没有安装依赖或执行训练。

import logging
import traceback

from datasets import load_dataset

from sentence_transformers.cross_encoder import (
    CrossEncoder,
    CrossEncoderModelCardData,
    CrossEncoderTrainer,
    CrossEncoderTrainingArguments,
)
from sentence_transformers.cross_encoder.evaluation import CrossEncoderNanoBEIREvaluator
from sentence_transformers.cross_encoder.losses import CachedMultipleNegativesRankingLoss

# 开启 INFO 日志,观察训练流程。
logging.basicConfig(
    format="%(asctime)s - %(message)s",
    datefmt="%Y-%m-%d %H:%M:%S",
    level=logging.INFO,
)

model_name = "microsoft/MiniLM-L12-H384-uncased"
train_batch_size = 64
num_epochs = 1
num_rand_negatives = 5  # 每个问答正例配用的随机负例数

# 1. 加载待微调模型和可选模型卡信息。
model = CrossEncoder(
    model_name,
    model_card_data=CrossEncoderModelCardData(
        language="en",
        license="apache-2.0",
        model_name="MiniLM-L12-H384 trained on GooAQ",
    ),
    model_kwargs={"torch_dtype": "float32"},
)
print("Model max length:", model.max_length)
print("Model num labels:", model.num_labels)

# 2. 取 GooAQ 的 100,000 条样本,固定随机种子拆出 1,000 条留作评估。
logging.info("Read the gooaq training dataset")
full_dataset = load_dataset(
    "sentence-transformers/gooaq", split="train"
).select(range(100_000))
dataset_dict = full_dataset.train_test_split(test_size=1_000, seed=12)
train_dataset = dataset_dict["train"]
eval_dataset = dataset_dict["test"]
logging.info(train_dataset)
logging.info(eval_dataset)

# 3. 为每个 query-answer 正例采样随机负例;mini_batch_size 影响显存用量。
loss = CachedMultipleNegativesRankingLoss(
    model=model,
    num_negatives=num_rand_negatives,
    mini_batch_size=32,
)

# 4. NanoBEIR 是轻量英文重排评估器;先评估未微调模型作为基线。
evaluator = CrossEncoderNanoBEIREvaluator(
    dataset_names=["msmarco", "nfcorpus", "nq"],
    batch_size=train_batch_size,
)
evaluator(model)

# 5. 设置训练参数。
short_model_name = (
    model_name if "/" not in model_name else model_name.split("/")[-1]
)
run_name = f"reranker-{short_model_name}-gooaq-cmnrl"
args = CrossEncoderTrainingArguments(
    output_dir=f"models/{run_name}",
    num_train_epochs=num_epochs,
    per_device_train_batch_size=train_batch_size,
    per_device_eval_batch_size=train_batch_size,
    learning_rate=2e-5,
    warmup_steps=0.1,
    fp16=False,  # 若当前 GPU 不支持 FP16,可保持关闭
    bf16=False,  # 上游设为 True;本文默认关闭,核实硬件和依赖后再启用
    eval_strategy="steps",
    eval_steps=0.1,
    save_strategy="steps",
    save_steps=0.1,
    save_total_limit=2,
    logging_steps=0.01,
    logging_first_step=True,
    run_name=run_name,
    report_to="none",  # 默认关闭 W&B 等远程指标上报
    seed=12,
)

# 6. 由 Trainer 连接模型、训练集、留出集、损失和评估器并开始训练。
trainer = CrossEncoderTrainer(
    model=model,
    args=args,
    train_dataset=train_dataset,
    eval_dataset=eval_dataset,
    loss=loss,
    evaluator=evaluator,
)
trainer.train()

# 7. 训练后再次运行 NanoBEIR 评估,便于将指标写入模型卡。
evaluator(model)

# 8. 将模型保存在本地。
final_output_dir = f"models/{run_name}/final"
model.save_pretrained(final_output_dir)

# 9. 上游示例还包含 model.push_to_hub(run_name) 及异常处理。
# 本文将远程上传整段保持注释,避免复制脚本后意外向 Hub 写入模型。
# 如要发布,先核对模型/训练数据各自许可、目标仓库和访问权限;
# 再在明确授权的环境中登录(当前 CLI 形式为 `hf auth login`)并手动启用。
# try:
#     model.push_to_hub(run_name)
# except Exception:
#     logging.error(
#         f"Error uploading model to the Hugging Face Hub:\n"
#         f"{traceback.format_exc()}"
#     )

原文简例用 microsoft/MiniLM-L12-H384-uncased、GooAQ 前十万问答(1,000 条留作验证)、批大小 64 和一个 epoch,配 CachedMultipleNegativesRankingLoss 的 5 个随机负例、mini_batch_size=32,并用 NanoBEIR 在训练前后评估,最后 save_pretrained。扩展示例改用 ModernBERT、5 个挖掘负例和二元交叉熵,组合任务内重排与 NanoBEIR,并按 eval_gooaq-dev_ndcg@10 保留最佳模型。

以下以扩展示例整理出完整主流程。与原文的差异:显式 num_labels=1;用当前明确参数 absolute_margin 替代旧 margin 别名;按实际负 / 正样本数计算 pos_weight;默认关闭混合精度与远程指标上报;只保存本地,不自动 push_to_hub。没有固定模型 revision 或软件版本,因此运行者仍应在自己的依赖锁定流程中补齐。本稿没有运行它。

from pathlib import Path
import logging
import traceback
import torch
from datasets import load_dataset
from sentence_transformers import SentenceTransformer
from sentence_transformers.cross_encoder import (
    CrossEncoder, CrossEncoderModelCardData, CrossEncoderTrainer, CrossEncoderTrainingArguments,
)
from sentence_transformers.cross_encoder.evaluation import (
    CrossEncoderNanoBEIREvaluator, CrossEncoderRerankingEvaluator,
)
from sentence_transformers.cross_encoder.losses import BinaryCrossEntropyLoss
from sentence_transformers.base.evaluation import SequentialEvaluator
from sentence_transformers.util import mine_hard_negatives

logging.basicConfig(
    format="%(asctime)s - %(message)s",
    datefmt="%Y-%m-%d %H:%M:%S",
    level=logging.INFO,
)


def main():
    output = Path("models/wwj-modernbert-gooaq")
    run_name = output.name  # 本地目录名也作为可选 Hub 仓库名;上传仍默认关闭
    if output.exists():
        raise FileExistsError("Choose a new output directory")

    model = CrossEncoder(
        "answerdotai/ModernBERT-base",
        num_labels=1,
        model_card_data=CrossEncoderModelCardData(
            language="en",
            model_name="ModernBERT-base trained on GooAQ",
            # 上游示例填了 license="apache-2.0";这里不预填,
            # 训练运行者应先核对基座模型、数据集和衍生权重许可。
        ),
        model_kwargs={"torch_dtype": "float32"},
    )
    logging.info("Read the GooAQ training dataset")
    full = load_dataset("sentence-transformers/gooaq", split="train")
    full = full.select(range(100_000)).select_columns(["question", "answer"])
    splits = full.train_test_split(test_size=1_000, seed=12)
    logging.info(splits["train"])
    logging.info(splits["test"])
    embedding = SentenceTransformer(
        "sentence-transformers/static-retrieval-mrl-en-v1", device="cpu"
    )
    hard_train = mine_hard_negatives(
        splits["train"], embedding,
        num_negatives=5, absolute_margin=0,
        range_min=0, range_max=100, sampling_strategy="top",
        batch_size=4096, output_format="labeled-pair", use_faiss=True,
    )
    logging.info(hard_train)
    labels = list(hard_train["label"])
    positive_count = labels.count(1)
    negative_count = labels.count(0)
    if not positive_count or not negative_count:
        raise ValueError("Training needs both positive and negative pairs")
    loss = BinaryCrossEntropyLoss(
        model=model,
        pos_weight=torch.tensor(negative_count / positive_count),
    )

    hard_eval = mine_hard_negatives(
        splits["test"], embedding,
        corpus=full["answer"], num_negatives=30,
        batch_size=4096, include_positives=True,
        output_format="n-tuple", use_faiss=True,
    )
    logging.info(hard_eval)
    samples = [
        {
            "query": row["question"],
            "positive": [row["answer"]],
            "documents": [row[name] for name in hard_eval.column_names[2:]],
        }
        for row in hard_eval
    ]
    rerank = CrossEncoderRerankingEvaluator(
        samples=samples, batch_size=64, name="gooaq-dev",
        always_rerank_positives=False,
    )
    nano = CrossEncoderNanoBEIREvaluator(
        dataset_names=["msmarco", "nfcorpus", "nq"], batch_size=64,
    )
    evaluator = SequentialEvaluator([rerank, nano])
    evaluator(model)  # 基础模型基线;调用时才真正执行评估。

    args = CrossEncoderTrainingArguments(
        output_dir=str(output),
        num_train_epochs=1,
        per_device_train_batch_size=64,
        per_device_eval_batch_size=64,
        learning_rate=2e-5, warmup_steps=0.1,
        fp16=False, bf16=False,
        dataloader_num_workers=2, dataloader_persistent_workers=True,
        load_best_model_at_end=True,
        metric_for_best_model="eval_gooaq-dev_ndcg@10",
        eval_strategy="steps", eval_steps=0.1,
        save_strategy="steps", save_steps=0.1, save_total_limit=2,
        logging_steps=0.01, logging_first_step=True,
        report_to="none", push_to_hub=False, seed=12,
    )
    trainer = CrossEncoderTrainer(
        model=model, args=args, train_dataset=hard_train,
        loss=loss, evaluator=evaluator,
    )
    trainer.train()
    evaluator(model)
    model.save_pretrained(str(output / "final"))
    logging.info("Saved the final model locally to %s", output / "final")

    # 上游示例含可选 model.push_to_hub(run_name) 和异常日志处理。
    # 默认关闭;只有在核对模型/数据许可、目标仓库与可见性,并由运行者明确选择后,
    # 才把下一行改为 True。认证建议使用当前 `hf auth login` 命令;本文未登录或上传。
    ENABLE_HUB_UPLOAD = False
    if ENABLE_HUB_UPLOAD:
        try:
            model.push_to_hub(run_name)
        except Exception:
            logging.error(
                "Error uploading model to the Hugging Face Hub:\n%s",
                traceback.format_exc(),
            )


if __name__ == "__main__":
    main()

代码边界:这份流程会联网下载公开模型与数据,并写入缓存和输出目录;十万条数据、FAISS 和批大小 64 都有实际资源需求。应使用隔离训练环境,先按硬件缩小样本和批量。输出目录存在时明确退出,避免不小心混入旧运行。仅按问题随机分割不能自动消除近重复或跨集合内容泄漏,真实任务应按数据来源、实体或时间进一步划分并去重。训练数据和公开评估基准的许可应单独核对,不能因代码库是 Apache-2.0 就推定所有数据和模型也是同一许可。

原文对 pos_weight 的注释把比例方向说得含混,示例实际使用 5 来提高较少的正例权重。这里明确采用负例数 / 正例数;挖掘后实际比例不一定恰好为 5。改写逻辑与公式已说明,但没有运行验证训练器的版本匹配或数值结果。

回调、多数据集和容易忽略的训练问题

Trainer 集成 Transformers 的回调,例如安装 wandb 时的 WandbCallback、可访问 tensorboard 时的 TensorBoardCallback,以及安装 codecarbon 后的 CodeCarbonCallback。后者记录的估算排放可进入自动模型卡。远程记录可能发送运行名称与指标,原文把 W&B 记录和 Hub 上传作为可用功能;本稿示例显式 report_to=”none”、push_to_hub=False,未经自己的选择不自动外发。

多数据集训练只需传入名称到 Dataset 的字典,或 DatasetDict;如需不同损失,再传同名损失字典。每个批次只来自一个数据集,不要求把所有数据都转换成同一格式。MultiDatasetBatchSamplers.ROUND_ROBIN 轮流取样,到任一数据集耗尽便停止,可能没用完其他数据,但各数据集被取样次数相同;默认 PROPORTIONAL 按数据量比例采样,会用完每个集合,大集合出现更频繁。

CrossEncoder 容易过拟合。原文建议结合 NanoBEIR 或任务内重排评估器,设置 load_best_model_at_end 和 metric_for_best_model。难负例能教模型区分“与问题相关”与“真正回答问题”,但只有难负例也可能削弱简单任务表现,甚至重排前 200 个候选的 top-10 比只重排前 100 个更差。混入随机负例可以缓解这一现象,具体比例必须通过自己的评估选择。

旧训练接口与当前接口

v4.0 之前常用 InputExample、DataLoader 和 CrossEncoder.fit():构造带标签文本对列表,放入 shuffle 的 DataLoader,再调用 fit 指定 epoch、warmup。两条示例文本只说明 API,远不足以训练。

v4.0 之后仍可调用 fit,但它会在内部初始化 CrossEncoderTrainer;直接用 Trainer 能通过 TrainingArguments 获得更细控制。旧行为还可通过 old_fit 获取,但原文说明未来计划完全弃用。新工作应优先使用当前 Trainer,不要把旧脚本参数与新接口混用。

v4.0 以前的完整旧版 fit 示例

以下是原文的旧接口片段,保留 InputExample、打乱样本的 DataLoader 和 CrossEncoder.fit() 调用。它演示 API 迁移,不是当前推荐的训练方案;两条样本仅是结构示例,不能构成有效训练集。

from sentence_transformers import CrossEncoder, InputExample
from torch.utils.data import DataLoader

# 定义 CrossEncoder;可以从头创建,或加载预训练模型。
model = CrossEncoder("distilbert/distilbert-base-uncased")

# 训练样本。生产任务需要远多于这里的两个演示样本。
train_examples = [
    InputExample(texts=["What are pandas?", "The giant panda ..."], label=1),
    InputExample(texts=["What's a panda?", "Mount Vesuvius is a ..."], label=0),
]

# 构造打乱顺序的 DataLoader,并训练一个 epoch。
train_dataloader = DataLoader(train_examples, shuffle=True, batch_size=16)
model.fit(train_dataloader=train_dataloader, epochs=1, warmup_steps=100)

原文说明:v4.0 起 fit() 仍可用,但内部改为初始化 CrossEncoderTrainer;遇到兼容问题时,old_fit() 可恢复旧行为,但计划完全弃用。新代码应直接使用 Trainer 与 TrainingArguments。本文没有运行旧 API,也未验证当前安装版本是否仍保留它。

与 SentenceTransformer 训练相比,一个关键差异是 CrossEncoder 的数据列可包含变长文本列表,例如 ListNetLoss 的需要;SentenceTransformer 相应训练接口不接受这种列形式。完整应用案例可继续阅读官方 Training and Finetuning Reranker Models 与多模态训练文章,它们是延伸阅读,本文没有将其未复现的结果作为本次证据。

训练结束可用 save_pretrained 保存本地。原文的可选上传使用了旧 huggingface-cli login 名称;当前官方 CLI 是 hf auth login。登录会保存凭据,push_to_hub 会向远程仓库上传;两者都不是本地训练必需步骤,本稿没有执行。没有发现静态问题不等于没有漏洞,也不等于模型已经适合生产检索。

原作者、出处与完整许可

出处:Cross Encoder Training Overview,Sentence Transformers 官方文档贡献者(原页无个人署名);持续更新的文档,全文快照核读于 2026-10-05,发布前于 2026-10-08 再次核验官网。文档源代码、许可证和归属说明见 Sentence Transformers 官方仓库。本文为未完纪独立中文翻译与编辑,并非官方发布或背书;保留原作者及项目归属。配图为未完纪原创技术示意图,非运行截图。转载与许可条件见附带文本;模型、数据集和第三方资产仍各有自己的许可,不由文档许可覆盖。

本文随文附上完整 Apache License 2.0 文本与上游 NOTICE 归属声明;可直接查阅 Sentence Transformers 官方许可证及 官方 NOTICE。本文为未完纪独立中文翻译与编辑,保留原作者和项目归属,不代表 Sentence Transformers 或其维护者的官方发布或背书。配图为未完纪原创技术示意图,非运行截图。模型、数据集、基准及第三方依赖各自适用独立许可,发布或再训练前应逐项核实。

Sentence Transformers 上游 NOTICE 原文
-------------------------------------------------------------------------------
Sentence Transformers
Copyright 2019-2025
Ubiquitous Knowledge Processing (UKP) Lab
Technische Universität Darmstadt
Copyright 2025-present
Hugging Face, Inc.
-------------------------------------------------------------------------------
Apache License 2.0 全文
                                Apache License
                           Version 2.0, January 2004
                        http://www.apache.org/licenses/

   TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION

   1. Definitions.

      "License" shall mean the terms and conditions for use, reproduction,
      and distribution as defined by Sections 1 through 9 of this document.

      "Licensor" shall mean the copyright owner or entity authorized by
      the copyright owner that is granting the License.

      "Legal Entity" shall mean the union of the acting entity and all
      other entities that control, are controlled by, or are under common
      control with that entity. For the purposes of this definition,
      "control" means (i) the power, direct or indirect, to cause the
      direction or management of such entity, whether by contract or
      otherwise, or (ii) ownership of fifty percent (50%) or more of the
      outstanding shares, or (iii) beneficial ownership of such entity.

      "You" (or "Your") shall mean an individual or Legal Entity
      exercising permissions granted by this License.

      "Source" form shall mean the preferred form for making modifications,
      including but not limited to software source code, documentation
      source, and configuration files.

      "Object" form shall mean any form resulting from mechanical
      transformation or translation of a Source form, including but
      not limited to compiled object code, generated documentation,
      and conversions to other media types.

      "Work" shall mean the work of authorship, whether in Source or
      Object form, made available under the License, as indicated by a
      copyright notice that is included in or attached to the work
      (an example is provided in the Appendix below).

      "Derivative Works" shall mean any work, whether in Source or Object
      form, that is based on (or derived from) the Work and for which the
      editorial revisions, annotations, elaborations, or other modifications
      represent, as a whole, an original work of authorship. For the purposes
      of this License, Derivative Works shall not include works that remain
      separable from, or merely link (or bind by name) to the interfaces of,
      the Work and Derivative Works thereof.

      "Contribution" shall mean any work of authorship, including
      the original version of the Work and any modifications or additions
      to that Work or Derivative Works thereof, that is intentionally
      submitted to Licensor for inclusion in the Work by the copyright owner
      or by an individual or Legal Entity authorized to submit on behalf of
      the copyright owner. For the purposes of this definition, "submitted"
      means any form of electronic, verbal, or written communication sent
      to the Licensor or its representatives, including but not limited to
      communication on electronic mailing lists, source code control systems,
      and issue tracking systems that are managed by, or on behalf of, the
      Licensor for the purpose of discussing and improving the Work, but
      excluding communication that is conspicuously marked or otherwise
      designated in writing by the copyright owner as "Not a Contribution."

      "Contributor" shall mean Licensor and any individual or Legal Entity
      on behalf of whom a Contribution has been received by Licensor and
      subsequently incorporated within the Work.

   2. Grant of Copyright License. Subject to the terms and conditions of
      this License, each Contributor hereby grants to You a perpetual,
      worldwide, non-exclusive, no-charge, royalty-free, irrevocable
      copyright license to reproduce, prepare Derivative Works of,
      publicly display, publicly perform, sublicense, and distribute the
      Work and such Derivative Works in Source or Object form.

   3. Grant of Patent License. Subject to the terms and conditions of
      this License, each Contributor hereby grants to You a perpetual,
      worldwide, non-exclusive, no-charge, royalty-free, irrevocable
      (except as stated in this section) patent license to make, have made,
      use, offer to sell, sell, import, and otherwise transfer the Work,
      where such license applies only to those patent claims licensable
      by such Contributor that are necessarily infringed by their
      Contribution(s) alone or by combination of their Contribution(s)
      with the Work to which such Contribution(s) was submitted. If You
      institute patent litigation against any entity (including a
      cross-claim or counterclaim in a lawsuit) alleging that the Work
      or a Contribution incorporated within the Work constitutes direct
      or contributory patent infringement, then any patent licenses
      granted to You under this License for that Work shall terminate
      as of the date such litigation is filed.

   4. Redistribution. You may reproduce and distribute copies of the
      Work or Derivative Works thereof in any medium, with or without
      modifications, and in Source or Object form, provided that You
      meet the following conditions:

      (a) You must give any other recipients of the Work or
          Derivative Works a copy of this License; and

      (b) You must cause any modified files to carry prominent notices
          stating that You changed the files; and

      (c) You must retain, in the Source form of any Derivative Works
          that You distribute, all copyright, patent, trademark, and
          attribution notices from the Source form of the Work,
          excluding those notices that do not pertain to any part of
          the Derivative Works; and

      (d) If the Work includes a "NOTICE" text file as part of its
          distribution, then any Derivative Works that You distribute must
          include a readable copy of the attribution notices contained
          within such NOTICE file, excluding those notices that do not
          pertain to any part of the Derivative Works, in at least one
          of the following places: within a NOTICE text file distributed
          as part of the Derivative Works; within the Source form or
          documentation, if provided along with the Derivative Works; or,
          within a display generated by the Derivative Works, if and
          wherever such third-party notices normally appear. The contents
          of the NOTICE file are for informational purposes only and
          do not modify the License. You may add Your own attribution
          notices within Derivative Works that You distribute, alongside
          or as an addendum to the NOTICE text from the Work, provided
          that such additional attribution notices cannot be construed
          as modifying the License.

      You may add Your own copyright statement to Your modifications and
      may provide additional or different license terms and conditions
      for use, reproduction, or distribution of Your modifications, or
      for any such Derivative Works as a whole, provided Your use,
      reproduction, and distribution of the Work otherwise complies with
      the conditions stated in this License.

   5. Submission of Contributions. Unless You explicitly state otherwise,
      any Contribution intentionally submitted for inclusion in the Work
      by You to the Licensor shall be under the terms and conditions of
      this License, without any additional terms or conditions.
      Notwithstanding the above, nothing herein shall supersede or modify
      the terms of any separate license agreement you may have executed
      with Licensor regarding such Contributions.

   6. Trademarks. This License does not grant permission to use the trade
      names, trademarks, service marks, or product names of the Licensor,
      except as required for reasonable and customary use in describing the
      origin of the Work and reproducing the content of the NOTICE file.

   7. Disclaimer of Warranty. Unless required by applicable law or
      agreed to in writing, Licensor provides the Work (and each
      Contributor provides its Contributions) on an "AS IS" BASIS,
      WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
      implied, including, without limitation, any warranties or conditions
      of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
      PARTICULAR PURPOSE. You are solely responsible for determining the
      appropriateness of using or redistributing the Work and assume any
      risks associated with Your exercise of permissions under this License.

   8. Limitation of Liability. In no event and under no legal theory,
      whether in tort (including negligence), contract, or otherwise,
      unless required by applicable law (such as deliberate and grossly
      negligent acts) or agreed to in writing, shall any Contributor be
      liable to You for damages, including any direct, indirect, special,
      incidental, or consequential damages of any character arising as a
      result of this License or out of the use or inability to use the
      Work (including but not limited to damages for loss of goodwill,
      work stoppage, computer failure or malfunction, or any and all
      other commercial damages or losses), even if such Contributor
      has been advised of the possibility of such damages.

   9. Accepting Warranty or Additional Liability. While redistributing
      the Work or Derivative Works thereof, You may choose to offer,
      and charge a fee for, acceptance of support, warranty, indemnity,
      or other liability obligations and/or rights consistent with this
      License. However, in accepting such obligations, You may act only
      on Your own behalf and on Your sole responsibility, not on behalf
      of any other Contributor, and only if You agree to indemnify,
      defend, and hold each Contributor harmless for any liability
      incurred by, or claims asserted against, such Contributor by reason
      of your accepting any such warranty or additional liability.

   END OF TERMS AND CONDITIONS

   APPENDIX: How to apply the Apache License to your work.

      To apply the Apache License to your work, attach the following
      boilerplate notice, with the fields enclosed by brackets "{}"
      replaced with your own identifying information. (Don't include
      the brackets!)  The text should be enclosed in the appropriate
      comment syntax for the file format. We also recommend that a
      file or class name and description of purpose be included on the
      same "printed page" as the copyright notice for easier
      identification within third-party archives.

   Copyright 2019 Nils Reimers

   Licensed under the Apache License, Version 2.0 (the "License");
   you may not use this file except in compliance with the License.
   You may obtain a copy of the License at

       http://www.apache.org/licenses/LICENSE-2.0

   Unless required by applicable law or agreed to in writing, software
   distributed under the License is distributed on an "AS IS" BASIS,
   WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
   See the License for the specific language governing permissions and
limitations under the License.
© 版权声明
THE END
喜欢就支持一下吧
点赞0 分享
评论 抢沙发

请登录后发表评论

    暂无评论内容