ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

sentence-transformers CrossEncoderTrainer 完全指南:基于 Transformers Trainer 的交叉编码器训练、评估与调参

sentence-transformers CrossEncoderTrainer 完全指南:基于 Transformers Trainer 的交叉编码器训练、评估与调参 人工智能NLPEmbedding微调【免费下载链接】sentence-transformersState-of-the-Art Embeddings, Retrieval, and Reranking项目地址https://gitcode.com/gh_mirrors/se/sentence-transformers点击查看免费下载导读CrossEncoderTrainer是 sentence-transformers 为交叉编码器CrossEncoder模型提供的官方训练器它构建在 Hugging Face TransformersTrainer之上将模型、数据集、损失函数、评估器与训练参数整合进一个统一、功能完备的训练与评估循环。本文围绕 CrossEncoderTrainer 参考文档结合 trainer.py 源码、BaseTrainer 与仓库内的真实示例和测试系统讲解其构造参数、默认行为、回调机制、数据整理流程与端到端实战用法。读完本文你将能够用CrossEncoderTrainer独立完成重排Reranking、语义文本相似度STS、二分类/多标签分类等交叉编码器任务的训练、评估、日志记录与模型保存。什么是 CrossEncoderTrainerCrossEncoderTrainer定义在 sentence_transformers/cross_encoder/trainer.py是BaseTrainer的直接子类。BaseTrainer又继承自 Transformers 的Trainer因此它天然具备 Transformers Trainer 的全部能力优化器与学习率调度、梯度累积、混合精度、断点续训、分布式训练DDP/DeepSpeed、Checkpoint 保存与加载、向 Hub 推送模型等。与 TransformersTrainer相比CrossEncoderTrainer针对 Sentence Transformers 生态做了以下关键适配损失函数优先训练循环不再由模型forward返回 loss而是由用户显式传入的loss默认BinaryCrossEntropyLoss或CrossEntropyLoss计算损失评估双轨制既支持通过eval_dataset计算评估 loss也支持通过evaluator计算更实用的指标如 NDCG10两者可同时使用自动模型卡生成训练结束后会自动生成带训练参数与评估结果的模型卡README数据集预处理自动为多数据集训练添加dataset_name列、校验保留列名并为prompt/task传递提供数据整理支持。类级定义来自源码model_class CrossEncoder model_card_data_class CrossEncoderModelCardData model_card_callback_class CrossEncoderModelCardCallback data_collator_class CrossEncoderDataCollator training_args_class CrossEncoderTrainingArguments这五个类属性决定了CrossEncoderTrainer默认绑定的模型、模型卡数据、模型卡回调、数据整理器与训练参数类型任何子类都可以通过覆写它们来定制行为。构造函数与全部参数详解CrossEncoderTrainer.__init__的完整签名如下trainer.pyCrossEncoderTrainer( modelNone, argsNone, train_datasetNone, eval_datasetNone, lossNone, evaluatorNone, data_collatorNone, processing_classNone, model_initNone, compute_metricsNone, callbacksNone, optimizers(None, None), optimizer_cls_and_kwargsNone, preprocess_logits_for_metricsNone, )各参数的含义与注意事项参数类型说明modelCrossEncoder要训练、评估或用于预测的模型。若未提供则必须传model_initargsCrossEncoderTrainingArguments训练参数。默认使用CrossEncoderTrainingArguments(output_dirtmp_trainer)即输出到当前目录下的tmp_trainer目录base/trainer.pytrain_datasetDataset/DatasetDict/dict[str, Dataset]训练数据集格式必须能被所选损失函数接受见 Training Overview Dataset Formateval_dataset同上评估数据集用于计算评估 losslossnn.Module/dict[str, nn.Module]/ 工厂函数 / 工厂函数字典训练损失。字典形式用于多数据集训练时按数据集名选取损失工厂函数形式主要用于超参优化HPO。未提供时自动选择model.num_labels 1用BinaryCrossEntropyLoss否则用CrossEntropyLoss见 get_default_lossevaluatorBaseEvaluator/list[BaseEvaluator]训练期间计算更有用指标的评估器可与eval_dataset搭配或单独使用传入列表时自动包装为SequentialEvaluator顺序执行base/trainer.pydata_collatorCrossEncoderDataCollator默认由模型、args与processing_class自动创建见下文“数据整理”一节processing_classtokenizer / image processor 等默认取自model.processor旧版本参数名tokenizer已废弃为processing_class通过deprecated_kwargs兼容model_initCallable[[], CrossEncoder]模型初始化工厂用于超参搜索场景与model二选一compute_metricsCallable[[EvalPrediction], dict]当前不受支持传入会触发警告建议改用evaluator或eval_datasetbase/trainer.pycallbackslist[TrainerCallback]自定义回调叠加在默认回调之上可用trainer.remove_callback()移除默认回调optimizers(Optimizer, LambdaLR)元组默认使用torch.optim.AdamWtransformers.get_linear_schedule_with_warmupoptimizer_cls_and_kwargs(优化器类, 关键字参数字典)自定义优化器类及其参数preprocess_logits_for_metrics可调用对象评估/预测前对 logits 的预处理函数依赖与前置条件使用CrossEncoderTrainer需要安装accelerate与datasets。如果缺失构造函数会抛出提示性的RuntimeError要求通过以下命令安装base/trainer.pypip install -U sentence-transformers[train]此外如果设置了args.eval_strategy为非no但既未提供eval_dataset也未提供evaluator构造函数会抛出更友好的 ValueError 提示base/trainer.py。IterableDataset 限制由于accelerate会拼接 IterableDataset 的批次并期望 data collator 返回纯张量字典而CrossEncoderDataCollator返回的是原始文本列分词在损失函数内完成因此CrossEncoderTrainer不支持 IterableDataset。若传入会在__init__阶段直接抛出ValueError并提示先转换为Dataset或DatasetDicttrainer.py。训练参数CrossEncoderTrainingArgumentsCrossEncoderTrainingArgumentstraining_args.py继承自BaseTrainingArguments后者又继承自 TransformersTrainingArguments因此除下述 Sentence Transformers 特有参数外所有 Transformers 标准参数learning_rate、per_device_train_batch_size、num_train_epochs、bf16/fp16、gradient_accumulation_steps、eval_strategy、save_strategy、load_best_model_at_end等均可用。关键参数必须提供output_dir (str)模型 checkpoint 的输出目录。Sentence Transformers 特有参数prompts (Union[str, Dict[str, str]])训练/评估/测试数据集中使用的 prompt。由于 CrossEncoder 会把多个列的输入合并成句对不支持按列配置 prompt只支持两种格式str对所有数据集统一使用一个 prompt如promptsSearch: Dict[str, str]按数据集名映射 prompt仅当数据集是DatasetDict或dict[str, Dataset]时可用如prompts{dataset_a: Search: , dataset_b: Retrieve: }。若在CrossEncoderDataCollator中发现 per-column 的 prompt 字典value 仍是 dict会直接抛错data_collator.py。batch_sampler批采样器取值来自BatchSamplers枚举如BATCH_SAMPLER、NO_DUPLICATES、NO_DUPLICATES_HASHED、GROUP_BY_LABEL也支持传入自定义DefaultBatchSampler子类或工厂函数。默认BatchSamplers.BATCH_SAMPLER。批采样器的具体实现见 base/sampler.py 相关源码 与BaseTrainer.get_batch_samplerbase/trainer.py。multi_dataset_batch_sampler多数据集训练时的批采样器取值来自MultiDatasetBatchSamplersROUND_ROBIN/PROPORTIONAL默认PROPORTIONAL按数据集大小成比例采样。router_mapping (Dict[str, str])数据集名到 Router 路由如slow、fast的映射。同样因为输入列会合并成句对只支持按数据集映射如{dataset_a: slow, dataset_b: fast}。per-column 形式会被 data collator 拒绝data_collator.py。若模型包含 Router 模块但未提供router_mapping训练器也会抛出提示错误base/trainer.py。learning_rate_mapping (Dict[str, float] | None)参数名正则到学习率的映射可为模型不同部分设置不同学习率。例如{SparseStaticEmbedding\.*: 1e-3}会为SparseStaticEmbedding模块单独设置 1e-3 的学习率。其底层实现位于BaseTrainer.get_optimizer_cls_and_kwargs先从损失含损失函数内的可训练参数中收集参数再把匹配正则的参数从默认优化器组中抽出、单独成组并赋予指定的lr与weight_decay若正则匹配不到任何参数会抛错base/trainer.py。其他值得注意的默认行为来自BaseTrainingArguments.__post_init__base/training_args.py自动设置prediction_loss_onlyTrue使预测/评估阶段只计算 loss关闭ddp_broadcast_buffers避免基于 BertModel 的模型在 DDP 训练时报 inplace 操作错误非分布式多卡训练会警告建议改用 DDPDDP 模式下若dataloader_drop_lastFalse会警告并自动置为True以防批次不均导致挂起warmup_ratio/warmup_steps在 Transformers v4/v5 之间做了兼容转换。数据整理CrossEncoderDataCollatorCrossEncoderTrainer默认使用 CrossEncoderDataCollator。它与双编码器 collator 的关键区别是返回原始文本列而不是 tokenized 张量因为 CrossEncoder 的损失函数会在内部调用model.preprocess完成分词。每个 batch 中还会携带解析好的prompt与task按数据集名解析见_resolve_scalar供损失函数转发给model.preprocess。collator 还负责提取标签列通过valid_label_columns依次匹配标量标签转成torch.Tensorlist/tuple 等集合类型则转成张量列表透传dataset_name用于多数据集 多损失训练其余列按原样打包成 batch。列顺序非常重要。由于多个文本列会被按顺序合并成句对如果数据集的列顺序是answer, question那么MultipleNegativesRankingLoss会把answer当作 anchor、question当作正样本意外地优化成“给定答案猜问题”。因此构造数据集时务必确认列顺序与你的语义目标一致data_collator.py。损失计算流程compute_loss 与 collect_featuresCrossEncoderTrainer覆写了compute_losstrainer.py其核心流程如下从输入中弹出dataset_name、prompt、task元数据调用collect_featurestrainer.py把 batch 拆成features各句子的输入与labelslabel列若self.loss是字典且 batch 带有dataset_name则按数据集名选取对应的损失函数若模型被包装DDP/compile 等而损失函数内部持有旧模型引用则通过override_model_in_loss注入当前包装模型仅在model self.model_wrapped时执行以loss_fn(features, labels, promptprompt, tasktask)调用损失若损失函数不接受prompt/task关键字会捕获TypeError降级为loss_fn(features, labels)并输出一次警告若损失返回的是字典多分量损失调用track_loss_components累积各分量用于日志再求和得到总损失。其中collect_features会识别以input_ids、sentence_embedding、pixel_values、input_features、input_values、pixel_values_videos等后缀结尾的列按前缀分组还原为每个句子独立的特征字典base/trainer.py这也使得该训练器天然兼容文本、图像、音频等多种模态输入。回调机制与自动模型卡由于继承自 TransformersTrainerCrossEncoderTrainer自动集成各类TrainerCallbackWandbCallback安装wandb后自动把训练指标记录到 WB且会自动设置环境变量WANDB_PROJECTsentence-transformers若未手动指定base/trainer.pyTensorBoardCallback可访问tensorboard时记录到 TensorBoardCodeCarbonCallback安装codecarbon时追踪训练碳排放且这些碳排放数据会写入自动生成的模型卡TrackioCallbackTransformers v4.54同样会自动设置TRACKIO_PROJECT环境变量。此外训练器在初始化时还会注册CrossEncoderModelCardCallbackadd_model_card_callbackbase/trainer.py自动追踪训练参数、最佳 checkpoint 步骤等信息用于模型卡生成。测试 test_model_card.py 验证了未初始化 Trainer 时保存模型会复用原模型卡一旦创建CrossEncoderTrainer就会生成新的模型卡。评估机制eval_dataset 与 evaluator 双轨并行CrossEncoderTrainer.evaluate的评估输出由两部分合并而成base/trainer.pyeval_dataset的评估 loss来自evaluation_loop的标准输出evaluator的指标在evaluation_loop返回后执行结果以metric_key_prefix_为前缀合并进output.metrics。例如使用CrossEncoderNanoBEIREvaluator时指标键形如eval_NanoBEIR_R100_mean_ndcg10。evaluation_loop还做了两项重要处理多数据集训练时只在第一个数据集上跑 evaluator若is_in_train且eval_dataset是DatasetDict则仅对第一个子数据集运行 evaluator并以前缀eval而非eval_name记录指标避免重复开销base/trainer.py分布式评估通过distributed_evaluation上下文管理器确保 evaluator 只在主进程执行再通过broadcast_object_list把结果广播给所有进程FSDP/DeepSpeed 场景除外。BaseTrainer还实现了_load_best_model在load_best_model_at_endTrue时自动从state.best_model_checkpoint恢复最优 checkpoint并把最优步骤写入模型卡数据。端到端实战MS MARCO 重排模型训练仓库中的 training_ms_marco_bce.py 是使用CrossEncoderTrainer训练重排模型的完整示例浓缩了上述全部概念from datasets import load_dataset from sentence_transformers.base.sampler import BatchSamplers from sentence_transformers.cross_encoder import CrossEncoder from sentence_transformers.cross_encoder.evaluation import CrossEncoderNanoBEIREvaluator from sentence_transformers.cross_encoder.losses import BinaryCrossEntropyLoss from sentence_transformers.cross_encoder.trainer import CrossEncoderTrainer from sentence_transformers.cross_encoder.training_args import CrossEncoderTrainingArguments # 1. 定义模型训练建议加载 fp32 torch.manual_seed(12) model CrossEncoder(answerdotai/ModernBERT-base, model_kwargs{torch_dtype: float32}) # 2. 加载并转换 MS MARCO 数据集 dataset load_dataset(microsoft/ms_marco, v1.1, splittrain) def bce_mapper(batch): queries, passages, labels [], [], [] for query, passages_info in zip(batch[query], batch[passages]): for idx, is_selected in enumerate(passages_info[is_selected]): queries.append(query) passages.append(passages_info[passage_text][idx]) labels.append(is_selected) return {query: queries, passage: passages, label: labels} dataset dataset.map(bce_mapper, batchedTrue, remove_columnsdataset.column_names) dataset dataset.train_test_split(test_size10_000) train_dataset, eval_dataset dataset[train], dataset[test] # 3. 定义损失 loss BinaryCrossEntropyLoss(model) # 4. 定义评估器轻量级英文重排评估 evaluator CrossEncoderNanoBEIREvaluator(dataset_names[msmarco, nfcorpus, nq], batch_size32) # 5. 定义训练参数 args CrossEncoderTrainingArguments( output_dirmodels/reranker-msmarco-v1.1-ModernBERT-base-bce, num_train_epochs1, per_device_train_batch_size32, per_device_eval_batch_size32, learning_rate2e-5, warmup_steps0.1, # 也支持 float 比例形式 bf16True, # GPU 支持 BF16 时开启 batch_samplerBatchSamplers.BATCH_SAMPLER, load_best_model_at_endTrue, metric_for_best_modeleval_NanoBEIR_R100_mean_ndcg10, eval_strategysteps, eval_steps4_000, save_strategysteps, save_steps4_000, save_total_limit2, logging_steps1_000, logging_first_stepTrue, seed12, ) # 6. 创建训练器并开始训练 trainer CrossEncoderTrainer( modelmodel, argsargs, train_datasettrain_dataset, eval_dataseteval_dataset, lossloss, evaluatorevaluator, ) trainer.train() # 7. 评估最终模型结果会写入模型卡 evaluator(model) # 8. 保存最终模型 model.save_pretrained(models/reranker-msmarco-v1.1-ModernBERT-base-bce/final)这段代码完整展示了模型初始化含torch_dtype与随机种子控制、数据集转换列顺序即句对顺序、损失选择、CrossEncoderNanoBEIREvaluator评估器、以eval_NanoBEIR_R100_mean_ndcg10为最优模型选择指标的训练参数配置以及训练后评估与保存的完整闭环。仓库中 ms_marco 目录 还提供了 BCE、ADRMSE、LambdaRank、ListNet、ListMLE、PListMLE、RankNet、CMNRL 等多种损失函数的等价训练脚本可对照学习。多数据集训练按数据集映射损失当loss传入dict[str, nn.Module]时train_dataset/eval_dataset必须是DatasetDict且字典的键必须与数据集键完全对应否则会抛出明确错误测试用例见 test_trainer.py损失为字典但train_dataset不是DatasetDict→ValueError: If the provided loss is a dict, then the train_dataset must be a DatasetDict.数据集键未全部出现在损失字典中 → 列出缺失键的ValueError。BaseTrainer.preprocess_dataset会在必要时为数据集惰性添加dataset_name列通过set_transform或mapbase/trainer.pycompute_loss据此选择对应的损失函数。同时Dataset中名为return_loss与dataset_name的列是保留列名会被validate_column_names拒绝。深入源码的几条线索如果你想进一步研究训练器的底层机制以下文件与函数值得继续阅读sentence_transformers/base/trainer.pyBaseTrainer全部实现包括get_data_collator、get_batch_sampler、get_multi_dataset_batch_sampler、_build_dataloader多数据集时使用ConcatDatasetProportionalBatchSampler/RoundRobinBatchSampler、get_optimizer_cls_and_kwargs把损失函数内的可训练参数纳入优化器、_save同时保存模型、processing class 与training_args.bin、_load_from_checkpoint通过_load_with_module_classes恢复自定义模块sentence_transformers/cross_encoder/data_collator.py句对数据的整理逻辑与 prompt/task 解析sentence_transformers/cross_encoder/lossesBinaryCrossEntropyLoss、CrossEntropyLoss、MultipleNegativesRankingLoss、ListNet、ListMLE等损失函数配合 损失参考文档 阅读sentence_transformers/cross_encoder/model.pyCrossEncoder模型本身num_labels、preprocess、predict等tests/cross_encoder/test_trainer.py多数据集错误路径、模型卡复用、trainer 全流程的测试用例docs/cross_encoder/training_overview.md数据集格式规范每种损失接受的列格式。总结CrossEncoderTrainer是 sentence-transformers 交叉编码器训练的推荐入口它以 TransformersTrainer为底座融入了 Sentence Transformers 的损失驱动训练范式、双轨评估体系、自动模型卡与多数据集支持同时通过CrossEncoderTrainingArguments暴露了prompts、batch_sampler、router_mapping、learning_rate_mapping等针对交叉编码器定制的参数。结合本文给出的 MS MARCO 实战示例与源码线索你可以在此基础上快速迁移到 STS、Quora 重复问题检测、NLI 乃至多模态任意到任意检索等任务并利用 WB/TensorBoard 回调与自动模型卡实现完整的实验追踪闭环。赞分享人工智能NLPEmbedding微调【免费下载链接】sentence-transformersState-of-the-Art Embeddings, Retrieval, and Reranking项目地址https://gitcode.com/gh_mirrors/se/sentence-transformers点击查看免费下载相关推荐MiniCPM-SALA 微调实战指南基于 Transformers Trainer 与 LLaMA-Factory 的全参、LoRA 与分布式训练MiniCPM SALA 微调实战指南基于 Transformers Trainer 与 LLaMA Factory 的全参、LoRA 与分布式训练 本篇指南人工智能大模型基础模型本地部署微调openBMBMiniCPM-SALA 微调实战指南基于 Transformers Trainer 与 LLaMA-Factory 的全参数 / LoRA / 多节点训练MiniCPM SALA 微调实战指南基于 Transformers Trainer 与 LLaMA Factory 的全参数 / LoRA / 多节点训练人工智能大模型基础模型本地部署微调openBMBsentence-transformers交叉编码器CrossEncoder架构与排序任务sentence transformers交叉编码器CrossEncoder架构与排序任务 1. 交叉编码器CrossEncoder核心概念 1.1 什么人工智能NLPEmbedding微调创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表