sentence-transformers 多向量编码器训练参数详解:MultiVectorEncoderTrainingArguments 配置指南
2026/9/20 10:10:27 网站建设 项目流程
  • 人工智能
  • NLP
  • Embedding
  • 微调

【免费下载链接】sentence-transformers

State-of-the-Art Embeddings, Retrieval, and Reranking

项目地址:https://gitcode.com/gh_mirrors/se/sentence-transformers
点击查看免费下载

本文以 docs/package_reference/multi_vector_encoder/training_args.md 为核心,完整讲解MultiVectorEncoderTrainingArguments的字段体系与使用姿势。它面向 ColBERT 风格多向量(late-interaction)模型的训练,是MultiVectorEncoderTrainer的默认参数类。读完本文,你将掌握max_length的两种传参形式与 query/document 任务分配机制、与query_expansion的相互作用,以及从BaseTrainingArguments与 TransformersTrainingArguments继承而来的全部关键训练选项,并能在 MS MARCO 知识蒸馏、MIRIAD 对比学习等真实训练脚本中直接落地配置。

一、它是什么:一条参数类,串起整个多向量训练

sentence-transformers的多向量编码器(MultiVectorEncoder)与常见的单向量SentenceTransformer不同:它对每个输入产出逐 token 的向量序列,打分时使用 MaxSim 晚期交互算子——对每个查询 token 取与文档 token 的最大相似度,再在查询 token 上求和。这种架构(ColBERT 风格)在召回与重排任务上表现突出,但训练时的 token 长度控制、query/document 前缀路由、批量采样等都与单向量模型差异显著。

MultiVectorEncoderTrainingArguments正是为此而生的训练参数类,定义于 sentence_transformers/multi_vector_encoder/training_args.py,通过 docs/package_reference/multi_vector_encoder/training_args.md 的 autodoc 暴露为公开 API。它的继承链为:

transformers.TrainingArguments └── sentence_transformers.base.training_args.BaseTrainingArguments └── MultiVectorEncoderTrainingArguments
  • TransformersTrainingArguments:提供output_dirlearning_rateper_device_train_batch_sizegradient_accumulation_stepsfp16/bf16eval_strategysave_strategy等通用训练参数;
  • BaseTrainingArguments(sentence_transformers/base/training_args.py):叠加 ST 特有的promptsrouter_mappinglearning_rate_mappingbatch_samplermulti_dataset_batch_sampler等参数;
  • MultiVectorEncoderTrainingArguments:在_VALID_DICT_FIELDS中追加max_length,并新增唯一专属字段max_length——一个控制"训练期 token 截断长度"的参数。

配套使用的MultiVectorEncoderTrainer(见 sentence_transformers/multi_vector_encoder/trainer.py)通过training_args_class = MultiVectorEncoderTrainingArguments将二者绑定,并在未显式传loss时默认使用MultiVectorMultipleNegativesRankingLoss

二、核心专属字段:max_length

max_lengthMultiVectorEncoderTrainingArguments相对基类新增的唯一字段,默认值为None。它的语义精确定义在类 docstring 中,值得逐句拆解。

2.1 它只作用于训练,不改变模型自身配置

Maximum token length applied when tokenizing training and evaluation-loss batches, without changing the model's own configuration: evaluators, inference, and the saved model keep the lengths configured on the model (set those via ``processor_kwargs`` globally or ``processing_kwargs`` per call).

关键结论:

  • 该长度只施加于训练批次与评估损失(evaluation-loss)批次的 tokenize 过程;
  • 评估器(evaluators)、推理(inference)、保存的模型仍然使用模型自身的长度配置;
  • 若要修改模型自身的长度,应通过模型加载时的processor_kwargs(全局)或单次调用时的processing_kwargs设置,而不是训练参数。

这一设计的价值在于:训练可以用更短的长度上限以换取速度与显存,而推理与导出模型保持完整能力。docstring 明确指出,"Tight training caps can be much faster and measurably stronger than uncapped training"(紧凑的训练长度上限往往比不设上限更快、效果可测地更强),并指向多向量 MS MARCO 知识蒸馏示例。

2.2 两种传参格式

  1. int:对所有列统一施加同一最大长度。例如max_length=180

  2. Dict[str, int]:按任务(task)分别指定,例如{"query": 32, "document": 180}。此时任务分配规则为:

    • 数据整理器(data collator)默认将第 0 列赋予"query"任务,其余列赋予"document"任务
    • router_mapping参数可以按列名覆盖这一默认分配。

    该默认分配逻辑在 sentence_transformers/multi_vector_encoder/data_collator.py 的_get_task_for_column中有直接实现:

    def _get_task_for_column(self, column_name: str, column_position: int, router_mapping: dict[str, str]) -> str: task = router_mapping.get(column_name) if task is None: task = "query" if column_position == 0 else "document" return task

    注意:分配只看列位置(position),不看列名(column names are not consulted),这与多向量损失函数的位置化约定(column 0 = query)保持一致。需要按列名定制时用router_mapping覆盖。

2.3 与 query_expansion 的相互作用

这是多向量模型特有的重要细节:

  • 当模型配置了固定宽度的query_expansionstrategy="fixed")时,查询(query)会忽略max_length的覆盖——因为扩展长度已经固定了查询的宽度;
  • strategy="min"时,max_length作为上限(ceiling)生效,但永远不会低于扩展长度

也就是说,max_length不会把查询截得比 query expansion 的下限还短。这一约束与底层Transformer模块中query_expansion的校验逻辑呼应(见 sentence_transformers/base/modules/transformer.py,其中校验了query_length不得小于query_expansion['length'],以及"固定扩展长度下查询不可再截短"的语义)。

2.4 源码中的实现方式

在 sentence_transformers/multi_vector_encoder/training_args.py 中:

_VALID_DICT_FIELDS = [*BaseTrainingArguments._VALID_DICT_FIELDS, "max_length"] max_length: Union[int, None, dict[str, int]] = field( default=None, metadata={ "help": "Maximum token length for training and evaluation-loss tokenization. Either 1) an int " "applied to every column, or 2) a mapping of tasks ('query', 'document') to lengths. Evaluators, " "inference, and the saved model keep the model's own configuration." }, )

max_length注册进_VALID_DICT_FIELDS意味着:当通过命令行或字符串传参时,该字段会被当作可 JSON 解析的 dict 处理(详见下文第五节)。

三、继承自 BaseTrainingArguments 的 ST 特有参数

BaseTrainingArguments在 TransformersTrainingArguments之上新增了以下 ST 专属参数,多向量训练中同样全部可用。

3.1 prompts:为各列指定提示前缀

用于为训练、评估、测试数据集的每一列指定 prompt,支持四种格式:

格式说明
str单个 prompt,应用于所有列
Dict[str, str]列名 → prompt 的映射
Dict[str, str](数据集维度)数据集名 → prompt(仅当数据集为DatasetDict或 dict 时)
Dict[str, Dict[str, str]]数据集名 → (列名 → prompt)

对 ColBERT 风格模型,典型用法是{"query": "[Q] ", "document": "[D] "}(或模型实际使用的前缀 token)。需要说明:如果传入了纯字符串且无法解析为 JSON,__post_init__会把它当作作用于所有列的单个 prompt(见 sentence_transformers/base/training_args.py 的__post_init__逻辑)。

3.2 router_mapping:列 → Router 路由

A mapping of dataset column names to Router routes, like "query" or "document".

两种格式:

  1. Dict[str, str]:列名 → 路由(如{"query": "query", "passage": "document"});
  2. Dict[str, Dict[str, str]]:数据集名 → (列名 → 路由),用于多数据集训练/评估。

它决定了每个数据集列由哪个 Router 子模块处理,同时覆盖 data collator 的默认任务分配(默认"第 0 列为 query、其余为 document")。

3.3 learning_rate_mapping:分模块学习率

A mapping of parameter name regular expressions to learning rates.

允许对模型不同部分设置不同学习率,例如{'SparseStaticEmbedding\.*': 1e-3}。适用于只想以不同速率微调模型特定子模块(如投影层、特定 Embedding 层)的场景。

3.4 batch_sampler 与 multi_dataset_batch_sampler

  • batch_sampler:默认BatchSamplers.BATCH_SAMPLER,可选值见sentence_transformers.base.sampler.BatchSamplers。多向量对比学习训练中常使用BatchSamplers.NO_DUPLICATES(见下文 MIRIAD 示例)来降低批内重复文档;
  • multi_dataset_batch_sampler:默认MultiDatasetBatchSamplers.PROPORTIONAL,控制多数据集训练时的按比例采样。

两者都支持传入字符串枚举值(构造时自动转换)或自定义可调用对象;to_dict()时会剔除可调用对象以便序列化。

3.5 warmup 兼容逻辑

BaseTrainingArguments显式定义了warmup_ratio并实现了跨 Transformers 版本的兼容:

  • Transformers v5+warmup_ratio已废弃,使用warmup_steps(可传 float 表示比例);
  • Transformers v4:支持warmup_ratio与整数warmup_steps;若向warmup_steps传入(0, 1)区间的 float,会将其自动转换为warmup_ratio

这正是仓库示例中warmup_steps=0.05("Warm up over the first 5% of training steps")这一写法的来源。

四、其他在__post_init__中被自动修正的行为

BaseTrainingArguments.__post_init__中还有几个值得了解的训练行为(sentence_transformers/base/training_args.py):

  • prediction_loss_only = TrueSentenceTransformerTrainer.compute_loss被重写为只计算预测损失,因此显式设置以避免额外计算;
  • ddp_broadcast_buffers = False:避免基于 BertModel 的模型在 DDP 训练时触发 inplace 操作导致的RuntimeError
  • 非分布式模式下提示 DataParallel(DP)慢于 DistributedDataParallel(DDP);
  • DDP 模式下若未设置dataloader_drop_last,会自动置为True以避免不均匀末批次的挂起问题;
  • dataloader_num_workers > 0且工作进程通过spawn启动时,提示设置dataloader_persistent_workers=True,否则每个 worker 都要重新 import sentence-transformers(耗时数秒),反而比dataloader_num_workers=0更慢。

五、dict 字段的字符串解析机制

_VALID_DICT_FIELDS追踪所有"允许以字符串形式传入 dict"的字段,目前包括:

accelerator_config, fsdp_config, deepspeed, gradient_checkpointing_kwargs, lr_scheduler_kwargs, learning_rate_mapping, prompts, router_mapping

MultiVectorEncoderTrainingArguments追加了max_length)。

__post_init__中,learning_rate_mappingrouter_mapping若为字符串则尝试json.loads,解析失败会抛出明确的ValueError;而prompts解析失败时会被宽容地当作单 prompt 字符串。这意味着这些参数既可以在 Python 中以 dict 传入,也可以从命令行以 JSON 字符串传入。

六、实战一:MS MARCO 知识蒸馏训练(max_length=180)

仓库中的 examples/multi_vector_encoder/training/msmarco/training_kd.py 是一个完整可运行的 ColBERT 蒸馏训练脚本。其核心配置如下:

args = MultiVectorEncoderTrainingArguments( # Required parameter: output_dir=f"models/{run_name}", # Optional training parameters: num_train_epochs=num_epochs, per_device_train_batch_size=train_batch_size, gradient_accumulation_steps=gradient_accumulation_steps, per_device_eval_batch_size=train_batch_size, learning_rate=learning_rate, max_length=180, # Cap training tokenization at 180 tokens, the query width floor stays 32 via the expansion warmup_steps=0.05, # Warm up over the first 5% of training steps fp16=False, # Set to False if you get an error that your GPU can't run on FP16 bf16=True, # Set to True if you have a GPU that supports BF16 load_best_model_at_end=True, metric_for_best_model="eval_NanoBEIR_mean_maxsim_ndcg@10", # Optional tracking/debugging parameters: eval_strategy="steps", # The NanoBEIR evaluator runs on its own datasets, so no eval_dataset is needed eval_steps=0.1, save_strategy="steps", save_steps=0.1, save_total_limit=2, logging_steps=0.01, run_name=run_name, # Will be used in W&B if `wandb` is installed seed=42, )

要点解读:

  • 训练集为(query_id, document_ids, scores)的知识蒸馏格式,通过resolve_ids将 ID 实时解析为文本(max_list_length=32控制每个 query 的负例文档数);
  • 损失为MultiVectorDistillKLDivLoss(model=model, temperature=0.25)——教师分数分布经温度锐化后与学生分数分布做 KL 散度;
  • 评估器MultiVectorNanoBEIREvaluator自行加载 NanoBEIR 数据集,因此eval_strategy="steps"时无需eval_dataset
  • max_length=180的意义:训练 tokenize 被截断到 180 token,而查询宽度下限(32)由 query expansion 兜底;模型推理/保存仍保持自身完整长度配置;
  • metric_for_best_model="eval_NanoBEIR_mean_maxsim_ndcg@10"表明该模型以MeanMaxSim(长度归一化的 MaxSim)作为评估指标,训练期打分与评估口径保持一致;
  • 混合精度:bf16=Truefp16=False,需 GPU 支持 BF16。

七、实战二:MIRIAD 医疗问答对比学习(max_length=1024)

examples/multi_vector_encoder/training/miriad/training_contrastive.py 展示了从零构建 ColBERT 模块序列并训练的场景。其参数配置:

args = MultiVectorEncoderTrainingArguments( output_dir=f"models/{run_name}", num_train_epochs=num_epochs, per_device_train_batch_size=train_batch_size, per_device_eval_batch_size=8, learning_rate=learning_rate, max_length=1024, # Cap training tokenization: passages average ~940 tokens, the model serves 8192 warmup_steps=0.05, fp16=False, bf16=True, batch_sampler=BatchSamplers.NO_DUPLICATES, load_best_model_at_end=True, metric_for_best_model="eval_miriad_eval_maxsim_ndcg@10", eval_strategy="steps", eval_steps=0.1, save_strategy="steps", save_steps=0.1, save_total_limit=2, logging_steps=0.005, run_name=run_name, seed=42, )

要点解读:

  • 模型由Transformer + Dense(128) + MultiVectorMask(skiplist_words=punctuation) + Normalize四个模块顺序组成(可通过MultiVectorEncoder(modules=[...])从零构建,详见 sentence_transformers/multi_vector_encoder/model.py 的_load_default_modules与示例脚本中的构造方式);
  • 损失为CachedMultiVectorMultipleNegativesRankingLoss(model=model, mini_batch_size=mini_batch_size)——缓存式大 batch InfoNCE 目标;
  • max_length=1024的典型场景:MIRIAD 段落平均约 940 token,训练截断到 1024 即可覆盖绝大多数样本;而模型在推理时支持 8192 token(即文档注释"the model serves 8192")。训练与推理长度解耦正是max_length的设计意图;
  • batch_sampler=BatchSamplers.NO_DUPLICATES:避免同一批次内出现重复文档,提升对比学习质量;
  • 训练中评估用 1000 条 query 的子采样(build_ir_evaluator(..., max_queries=1000)),最终用完整 NanoBEIR 测试集评估。

八、常见的完整训练流程闭环

综合两个实战脚本,一个标准的多向量训练流程为:

  1. 构建/加载模型MultiVectorEncoder("lightonai/LateOn")MultiVectorEncoder(modules=[...])从零构建(投影层随机初始化时会有日志提示 "Training is required before this model is useful");
  2. 准备数据:标准 pair / triplet / multi-negative 格式,或(query, document_1, ..., document_N, scores)的蒸馏格式,配合resolve_ids实时解析;
  3. 选择损失:不传时默认MultiVectorMultipleNegativesRankingLoss;蒸馏用MultiVectorDistillKLDivLoss、margin-MSE 用MultiVectorMarginMSELoss,大 batch 用缓存式损失(可传loss=或按数据集名分发的 dict);
  4. 定义评估器:如MultiVectorNanoBEIREvaluatorMultiVectorInformationRetrievalEvaluator等(见 docs/package_reference/multi_vector_encoder/evaluation.md),可先跑一次基线;
  5. 配置MultiVectorEncoderTrainingArguments:重点调好max_length(考虑 query_expansion 的固定宽度下限)、metric_for_best_model与模型similarity_fn_name保持一致(maxsimmeanmaxsim)、batch sampler 与混合精度;
  6. 训练并保存MultiVectorEncoderTrainer(model=..., args=..., train_dataset=..., loss=..., evaluator=...)trainer.train()model.save_pretrained()model.push_to_hub()

九、关键设计要点速查

关注点结论依据
max_length作用域仅训练与 evaluation-loss 批次;评估器/推理/保存模型不受影响training_args.py docstring
max_length=int所有列统一长度同上
max_length=dict"query"/"document"任务区分;第 0 列默认 query,其余 documentdata_collator.py_get_task_for_column
任务覆盖router_mapping按列名覆盖默认分配同上 + base/training_args.py
query_expansion=fixed查询忽略max_length,宽度由扩展长度固定training_args.py docstring
query_expansion=minmax_length为上限,且不低于扩展长度training_args.py docstring
默认损失MultiVectorMultipleNegativesRankingLosstrainer.pyget_default_loss
推理长度配置走模型自身的processor_kwargs/processing_kwargstraining_args.py docstring
dict 字段字符串解析_VALID_DICT_FIELDS中的字段支持 JSON 字符串;max_length已追加training_args.py__post_init__

十、适用前提与注意事项

  • 本文描述的行为以当前仓库代码为准MultiVectorEncoderTrainingArguments是 2025 年后引入的新 API,若你使用的旧版本 sbert 中不存在MultiVectorEncoder或该参数类,请先升级到包含sentence_transformers/multi_vector_encoder/包的最新版本;
  • max_length的收益依赖模型本身支持的长序列能力与训练数据分布:文档与示例均强调"紧凑上限可能更快且更强",但具体数值(180、1024)应依据你自己的语料 token 分布决定;
  • 混合精度(fp16/bf16)需硬件支持;warmup_steps传 float 比例的写法依赖BaseTrainingArguments的 Transformers v4/v5 兼容逻辑;
  • 训练长度与推理长度解耦的前提是训练数据被截断后仍保留足够语义——对于超长文档的领域,过小的max_length会让模型"看不见"关键内容,建议先统计语料 token 分布再定值。

十一、继续深入阅读

  • 参数类实现:sentence_transformers/multi_vector_encoder/training_args.py、sentence_transformers/base/training_args.py
  • 配套 Trainer:sentence_transformers/multi_vector_encoder/trainer.py(默认损失、模型卡回调、混合精度、多 GPU)
  • 任务分配实现:sentence_transformers/multi_vector_encoder/data_collator.py
  • 模型与推理侧长度配置:sentence_transformers/multi_vector_encoder/model.py、sentence_transformers/base/modules/transformer.py(query_lengthdocument_lengthquery_expansion
  • 完整可运行示例:examples/multi_vector_encoder/training/msmarco/training_kd.py、examples/multi_vector_encoder/training/miriad/training_contrastive.py
  • 相关 API 参考:docs/package_reference/multi_vector_encoder/trainer.md、docs/package_reference/multi_vector_encoder/model.md、docs/package_reference/multi_vector_encoder/losses.md
  • 人工智能
  • NLP
  • Embedding
  • 微调

【免费下载链接】sentence-transformers

State-of-the-Art Embeddings, Retrieval, and Reranking

项目地址:https://gitcode.com/gh_mirrors/se/sentence-transformers
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询