Hugging Face发布完整教程,覆盖多模态嵌入和重排序模型的训练微调,可直接优化RAG系统
作为一个实际示例,我将为你详解如何对 Qwen/Qwen3-VL-Embedding-2B 进行微调,以完成视觉文档检索(VDR)任务——这是一个检索与给定文本查询相关的文档页面(作为图像,保留图表、表格和布局)的任务。我微调后的 tomaarsen/Qwen3-VL-Embedding-2B-vdr 模型展示了在自己的领域数据上微调能带来多少性能提升。在我的评估数据上,微调模型的 NDCG@10 达到 0.947,而基础模型只有 0.888,并且超过了我测试过的所有现有 VDR 模型,包括大小达 4 倍的模型。
如果你是 Sentence Transformers 多模态模型的新手,建议先阅读《使用 Sentence Transformers 进行多模态嵌入与排序模型》。要了解文本专用的嵌入、排序或稀疏嵌入模型的训练,见末尾的"先前博文"部分。
Visual Document Retrieval 数据集
CachedMultipleNegativesRankingLoss MatryoshkaLoss
CachedMultipleNegativesRankingLoss
模型大小 vs NDCG@10 Matryoshka 维度 vs NDCG@10
模型大小 vs NDCG@10
Matryoshka 维度 vs NDCG@10
先前博文 训练示例 文档
通用多模态嵌入模型(如 Qwen/Qwen3-VL-Embedding-2B)是在多样化数据上训练的,以在广泛的语言和任务上表现良好:图像-文本匹配、视觉问答、文档理解等。但这种通用性意味着该模型很少是任何特定任务的最佳选择。
考虑视觉文档检索:给定一个文本查询,如"该公司第三季度的收入是多少?",模型必须从数千份文档中找到最相关的文档截图。这需要理解文档布局、图表、表格和文本,这与例如将鞋子的图片与产品描述相匹配是截然不同的技能。
通过在特定领域数据上微调,模型可以学到这些专门的模式。在我的实验中,微调将 NDCG@10 从 0.888 提升到 0.947,超过了我测试过的每一个最近的多模态模型,包括大小达 4 倍的模型。
训练多模态 Sentence Transformer 模型涉及的组件与训练文本专用模型的组件相同:
模型:要训练或微调的多模态模型。
数据集:用于训练和评估的数据。
损失函数:量化模型性能并引导优化过程的函数。
训练参数(可选):影响训练性能和追踪/调试的参数。
评估器(可选):用于在训练前、训练期间或训练后评估模型的工具。
Trainer:将模型、数据集、损失函数和其他组件结合在一起进行训练。
多模态训练管道使用与文本专用训练相同的 SentenceTransformerTrainer。关键区别在于你的数据集包含图像(或其他模态)和文本,而模型的处理器会自动处理图像预处理。
让我们逐个讲解每个组件,以视觉文档检索(将文本查询与文档截图匹配)作为贯穿示例。
最常见的方法是微调现有的多模态嵌入模型,或从视觉-语言模型(VLM)检查点开始。Transformer 模块会从模型的处理器自动检测支持的模态。
要微调现有的多模态嵌入模型(例如已有 modules.json 文件的模型),你可以传入 processor_kwargs 和 model_kwargs 分别控制预处理和模型加载。processor_kwargs 直接传入 AutoProcessor.from_pretrained(...) (例如,图像分辨率边界:更高的 max_pixels 意味着更高质量但更多内存),而 model_kwargs 传入相应的 AutoModel.from_pretrained(...) 调用(例如,精度、注意力实现):
from sentence_transformers import SentenceTransformer
model = SentenceTransformer(
"Qwen/Qwen3-VL-Embedding-2B",
model_kwargs={"attn_implementation": "flash_attention_2", "torch_dtype": "bfloat16"},
processor_kwargs={"min_pixels": 28 * 28, "max_pixels": 600 * 600},
)
你也可以从一个还没有为嵌入训练过的全新 VLM 检查点开始。Sentence Transformers 会尝试识别架构,从处理器推断支持的模态,并设置相应的前向方法和池化。如果自动检测对某个特定模型不够完美,保存的 sentence_bert_config.json 中的配置可以编辑来调整模态设置、前向方法和输出处理:
from sentence_transformers import SentenceTransformer
model = SentenceTransformer("Qwen/Qwen3-VL-2B")
无论哪种情况,Transformer 模块都会检查处理器以确定哪些模态可用,并在必要时自动添加 Pooling。你可以验证支持的模态:
print(model.modalities)
# ['text', 'image', 'video', 'message']
print(model.supports("image"))
# True
你也可以使用 Router 模块为不同的模态组合单独的编码器,而不是使用单个 VLM 主干。这让你可以组合任何现有的编码器,并根据检测到的模态将输入路由到相应的编码器:
from sentence_transformers import SentenceTransformer
from sentence_transformers.sentence_transformer.modules import Dense, Pooling, Router, Transformer
# 为不同模态创建单独的编码器
text_encoder = Transformer("sentence-transformers/all-MiniLM-L6-v2")
text_pooling = Pooling(text_encoder.get_embedding_dimension(), pooling_mode="mean")
text_projection = Dense(text_encoder.get_embedding_dimension(), 768)
# SigLIP 直接输出池化嵌入,所以不需要单独的 Pooling 模块
image_encoder = Transformer("google/siglip2-base-patch16-224")
# 根据模态路由输入
router = Router(
sub_modules={
"text": [text_encoder, text_pooling, text_projection],
"image": [image_encoder],
},
)
model = SentenceTransformer(modules=[router])
由于基于 Router 的多模态模型为每个模态使用单独的编码器,它们的嵌入空间最初是未对齐的。需要训练来对齐这些空间以实现有意义的跨模态相似性。上面显示的 Dense 投影层有助于将来自不同编码器的嵌入映射到共享空间。
这种方法在你想使用轻量级的专用编码器而不是大型 VLM 时很有用。你还可以结合 Router 基础的多模态与基于任务的路由(例如为查询和文档使用不同的编码器)使用 route_mappings。更多高级路由场景见 Router 文档。
在这个示例中,我使用 tomaarsen/llamaindex-vdr-en-train-preprocessed 数据集,这是一个预处理的英文子集来自 llamaindex/vdr-multilingual-train。源数据集随《视觉文档检索走向多语言》博文由 LlamaIndex 发布,包含约 50 万个多语言查询-图像样本,来自公网 PDF,查询使用 VLM(gemini-1.5-pro 和 Qwen2-VL-72B)合成生成。我的预处理版本筛选到 53,512 个英文样本,并将每个样本的 16 个基于 ID 的困难负样本中的 4 个解析为实际文档截图图像,因此可以直接用于训练而无需进一步预处理:
from datasets import load_dataset
train_dataset = load_dataset("tomaarsen/llamaindex-vdr-en-train-preprocessed", "train", split="train")
train_dataset = train_dataset.select_columns(["query", "image", "negative_0"])
eval_dataset = load_dataset("tomaarsen/llamaindex-vdr-en-train-preprocessed", "eval", split="train")
train 配置包含前 10,000 个样本,eval 配置包含接下来的 300 个样本(一个包含全部 53,512 个样本的完整配置也可用)。对于训练,我选择 query、image 和 negative_0 来形成(锚点、正样本、困难负样本)三元组。包含额外的困难负样本可能会改进训练信号,但每个额外的负样本也会增加内存使用和训练时间,所以我只用一个。对于评估,我保留每个查询的所有四个困难负样本来构建一个更具挑战性的检索库(更多见下面的评估器部分)。
就像文本专用训练一样,数据集格式必须与你选择的损失函数相匹配。规则是一样的:
如果你的损失函数需要一个 Label,你的数据集必须有一个名为"label"或"score"的列。
除"label"或"score"外的所有列都被视为输入。这些列的数量必须与你选择的损失函数的有效输入数量相匹配。超过标签列之外,列名不重要,只有顺序重要。
对于多模态数据集,输入可以包含:
图像:PIL 图像、文件路径、URL 或 numpy/torch 数组。
音频:文件路径、numpy/torch 数组、带 "array" 和 "sampling_rate" 键的字典,或(如果安装了 torchcodec)torchcodec.AudioDecoder 实例。
视频:文件路径、numpy/torch 数组、带 "array" 和 "video_metadata" 键的字典,或(如果安装了 torchcodec)torchcodec.VideoDecoder 实例。
多模态字典:一个将模态名称映射到值的字典,例如 {"text": ..., "image": ...}。键必须是 "text"、"image"、"audio" 或 "video"。
数据整理器会自动调用 model.preprocess(),它检测每个输入的模态并应用相应的预处理。无需手动标记化或图像处理。
许多与 Sentence Transformers 开箱即用的 Hugging Face 数据集已标记为 sentence-transformers,使你可以轻松在 https://huggingface.co/datasets?other=sentence-transformers 找到它们。
在这次训练中,我使用 CachedMultipleNegativesRankingLoss,这是检索任务的常见选择。它接受 (query, positive) 对和任意数量的额外硬负样本列,从 0 到 n,只要每个样本具有相同数量的负样本即可。在训练过程中,损失函数将每个查询与其正样本的相似度推高,同时将其与每个负样本的相似度推低。负样本来自两个来源:
硬负样本:数据集中显式提供的负样本列(在我们的三元组设置中仅为 negative_0)。
批内负样本:来自同一批中其他每个样本的正样本和硬负样本,无额外成本地重新用作该查询的额外负样本。
每个查询的负样本越多意味着更强的训练信号,因此更大的批大小直接改进训练质量。除此之外,损失函数的"缓存"变体使用梯度缓存,即使 GPU 内存受限也能实现大的有效批大小。
mini_batch_size 参数控制在缓存前向传递期间一次处理多少个样本。对于大型多模态模型,将其设置为较小的值(例如 1)对于避免内存不足错误而不牺牲大有效批大小的好处很重要:
from sentence_transformers.sentence_transformer.losses import CachedMultipleNegativesRankingLoss
loss = CachedMultipleNegativesRankingLoss(model, mini_batch_size=1)
为了产生在多个维度上都能很好工作的嵌入,我使用 MatryoshkaLoss 包装基础损失。这训练模型,使得将嵌入截断到较小的维度数仍然能产生良好性能:
from sentence_transformers.sentence_transformer.losses import CachedMultipleNegativesRankingLoss, MatryoshkaLoss
loss = CachedMultipleNegativesRankingLoss(model, mini_batch_size=1)
loss = MatryoshkaLoss(model, loss, matryoshka_dims=[2048, 1536, 1024, 512, 256, 128, 64])
这对多模态模型特别有用,其中嵌入可能很大(Qwen3-VL 为 2048 维)。通过 Matryoshka 训练,你可以在部署时使用截断的嵌入(例如 256 或 128 维)以实现更快的搜索,同时质量损失最小。正如我将在结果部分展示的那样,微调模型即使在 512 维时也能达到接近峰值的性能。
SentenceTransformerTrainingArguments 类让你控制训练超参数。以下是用于 VDR 微调的配置:
from sentence_transformers.sentence_transformer.training_args import SentenceTransformerTrainingArguments, BatchSamplers
run_name = "Qwen3-VL-Embedding-2B-vdr"
args = SentenceTransformerTrainingArguments(
# Required parameter:
output_dir=f"models/{run_name}",
# Optional training parameters:
num_train_epochs=1,
per_device_train_batch_size=64,
per_device_eval_batch_size=64,
learning_rate=2e-5,
warmup_ratio=0.1,
fp16=False,
bf16=True,
batch_sampler=BatchSamplers.NO_DUPLICATES,
# Optional tracking/debugging parameters:
eval_strategy="steps",
eval_steps=0.1,
save_strategy="steps",
save_steps=0.1,
save_total_limit=2,
logging_steps=0.05,
run_name=run_name,
)
关于(多模态)训练需要注意的几点:
bf16=True:由于数值稳定性更好,通常优先选择 bfloat16 而不是 float16。
batch_sampler=BatchSamplers.NO_DUPLICATES:使用 MultipleNegativesRankingLoss 或其缓存变体时,批中没有重复样本可确保每个批内负样本都是真正不同的样本。
per_device_train_batch_size=64:对于 2B 参数的 VLM 来说这似乎很大,但 CachedMultipleNegativesRankingLoss 配合 mini_batch_size=1 通过梯度缓存处理内存限制。
eval_steps、save_steps 和 logging_steps:将这些设置为分数(例如 0.1)意味着评估、保存和日志记录将在每个 epoch 的 10% 时发生,这对于监控训练进度很有用。
为了在训练前、训练中和训练后跟踪检索性能,我使用 InformationRetrievalEvaluator。它计算标准检索指标,如 NDCG@10、MAP 和 Recall@k:
from sentence_transformers.sentence_transformer.evaluation import InformationRetrievalEvaluator
# Build the evaluation data from the eval dataset.
# Queries and corpus use integer IDs: query 0's relevant document is corpus 0.
eval_queries = {qid: sample["query"] for qid, sample in enumerate(eval_dataset)}
eval_corpus = {did: sample["image"] for did, sample in enumerate(eval_dataset)}
num_eval = len(eval_dataset)
# Add hard negatives to the corpus with offset IDs (num_eval, 2*num_eval, ...)
# so they don't collide with the positive document IDs (0..num_eval-1).
negative_columns = ["negative_0", "negative_1", "negative_2", "negative_3"]
for neg_idx, neg_col in enumerate(negative_columns):
for did, sample in enumerate(eval_dataset):
eval_corpus[num_eval * (neg_idx + 1) + did] = sample[neg_col]
# Each query's relevant document is the positive at the same index
eval_relevant_docs = {idx: [idx] for idx in range(len(eval_dataset))}
eval_evaluator = InformationRetrievalEvaluator(
queries=eval_queries,
corpus=eval_corpus,
relevant_docs=eval_relevant_docs,
batch_size=1,
show_progress_bar=True,
name="vdr-eval-hard",
)
评估器接受文本查询、图像语料库(包括硬负样本)以及文档与查询相关性的映射。注意语料库包含正样本和硬负样本文档截图的混合,使这成为具有挑战性的评估。使用 batch_size=1 可以防止在评估大型 VLM 期间出现内存不足问题。
SentenceTransformerTrainer 将所有内容整合在一起。以下是完整的训练脚本:
from datasets import load_dataset
from sentence_transformers import SentenceTransformer
from sentence_transformers.sentence_transformer.evaluation import InformationRetrievalEvaluator
from sentence_transformers.sentence_transformer.losses import CachedMultipleNegativesRankingLoss, MatryoshkaLoss
from sentence_transformers.sentence_transformer.model_card import SentenceTransformerModelCardData
from sentence_transformers.sentence_transformer.trainer import SentenceTransformerTrainer
from sentence_transformers.sentence_transformer.training_args import (
BatchSamplers,
SentenceTransformerTrainingArguments,
)
# 1. Load a model to finetune with (optional) model card data
model = SentenceTransformer(
"Qwen/Qwen3-VL-Embedding-2B",
model_card_data=SentenceTransformerModelCardData(
language="en",
license="apache-2.0",
model_name="Qwen3-VL-Embedding-2B model trained on Visual Document Retrieval query-document screenshot pairs",
),
model_kwargs={"attn_implementation": "flash_attention_2", "torch_dtype": "bfloat16"},
# Control image resolution: lower values save memory, higher values preserve detail
processor_kwargs={"min_pixels": 28 * 28, "max_pixels": 600 * 600},
)
# 2. Load a dataset to finetune on: (query, positive, negative_0) triplets for training,
# all 4 hard negatives retained for evaluation
train_dataset = load_dataset("tomaarsen/llamaindex-vdr-en-train-preprocessed", "train", split="train")
train_dataset = train_dataset.select_columns(["query", "image", "negative_0"])
eval_dataset = load_dataset("tomaarsen/llamaindex-vdr-en-train-preprocessed", "eval", split="train")
# 3. Define a loss function
loss = CachedMultipleNegativesRankingLoss(model, mini_batch_size=1)
loss = MatryoshkaLoss(model, loss, matryoshka_dims=[2048, 1536, 1024, 512, 256, 128, 64])
run_name = "Qwen3-VL-Embedding-2B-vdr"
args = SentenceTransformerTrainingArguments(
# 必需参数:
output_dir=f"models/{run_name}",
# 可选的训练参数:
num_train_epochs=1,
per_device_train_batch_size=64,
per_device_eval_batch_size=64,
learning_rate=2e-5,
warmup_ratio=0.1,
fp16=False, # BF16 相比 FP16 对 VLM 更优,数值稳定性更好
bf16=True, # 如果你的 GPU 支持 BF16,设置为 True(大多数现代 GPU 都支持)
batch_sampler=BatchSamplers.NO_DUPLICATES, # MultipleNegativesRankingLoss 受益于无重复数据
# 可选的跟踪/调试参数:
eval_strategy="steps",
eval_steps=0.1,
save_strategy="steps",
save_steps=0.1,
save_total_limit=2,
logging_steps=0.05,
run_name=run_name, # 用于 Trackio 等工具
# report_to=["codecarbon", "trackio"], # 取消注释以启用日志记录 (pip install codecarbon trackio)
)
eval_queries = {qid: sample["query"] for qid, sample in enumerate(eval_dataset)}
eval_corpus = {did: sample["image"] for did, sample in enumerate(eval_dataset)}
num_eval = len(eval_dataset)
negative_columns = ["negative_0", "negative_1", "negative_2", "negative_3"]
for neg_idx, neg_col in enumerate(negative_columns):
for did, sample in enumerate(eval_dataset):
eval_corpus[num_eval * (neg_idx + 1) + did] = sample[neg_col]
eval_relevant_docs = {idx: [idx] for idx in range(len(eval_dataset))}
eval_evaluator = InformationRetrievalEvaluator(
queries=eval_queries,
corpus=eval_corpus,
relevant_docs=eval_relevant_docs,
batch_size=1,
show_progress_bar=True,
name="vdr-eval-hard",
)
eval_evaluator(model)
trainer = SentenceTransformerTrainer(
model=model,
args=args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
loss=loss,
evaluator=eval_evaluator,
)
trainer.train()
eval_evaluator(model)
for dim in [2048, 1536, 1024, 512, 256, 128, 64]:
训练脚本几乎与纯文本训练脚本相同。唯一的区别是:
模型加载:我们传递 model_kwargs 用于精度和注意力实现,以及 processor_kwargs 用于图像分辨率边界。
损失函数:我们使用 CachedMultipleNegativesRankingLoss 与 mini_batch_size=1 来处理大型 VLM,避免内存溢出。
评估器:评估器在语料库中使用图像,文本作为查询,实现跨模态检索评估。
其他所有内容(训练器、训练参数、数据集加载)的工作方式与纯文本训练完全相同。
经过仅 1 个 epoch 的训练,微调后的 tomaarsen/Qwen3-VL-Embedding-2B-vdr 模型在评估集上达到 0.947 的 NDCG@10(300 个查询,1500 个语料库文档,余弦相似度)。这相比基础的 Qwen/Qwen3-VL-Embedding-2B 模型的 0.888 有了显著提升,并且优于所有现有的 VDR 模型:
微调后的 2B 模型甚至优于 8B 的 Qwen3-VL-Embedding 模型,展示了任务特定微调的强大力量。在你自己的领域上进行微调通常值得考虑,即使有更大的通用模型可用!
上面的比较使用了完整的 2048 维嵌入。得益于 Matryoshka 训练,微调后的模型在被截断到更少维度时也表现良好,让你在部署时可以在嵌入大小和检索质量之间进行权衡:
微调后模型在完整的 2048 维时达到峰值(0.948),但即使在 512 维时(缩小 4 倍)也能保持在峰值的 0.3% 以内,即使在 64 维时(缩小 32 倍)仍保留超过 92% 的峰值性能。Matryoshka 训练将最重要的信息集中在前面的维度中,因此适度的截断性能成本很小。
1024 维和 2048 维之间的差距很小(0.946 vs 0.948),所以我已将模型保存为在配置中设置 truncate_dim=1024。这意味着 SentenceTransformer("tomaarsen/Qwen3-VL-Embedding-2B-vdr") 默认会生成 1024 维嵌入,相比完整的 2048 维减半存储占用。如果你需要不同的维度,在加载时传递 truncate_dim=N 来覆盖它。
你也可以使用相同的训练基础设施来微调多模态 Cross Encoder(重排序)模型。关键区别是使用 CrossEncoderTrainer 和 Cross Encoder 特定的损失函数。本节提供简要概览;详细的完整、可运行的脚本包含数据集准备和评估,请参见完整训练示例。
以下是基于涂鸦训练脚本的简化示例,该脚本训练一个重排序器来匹配图像与文本标题:
from sentence_transformers.cross_encoder import CrossEncoder
from sentence_transformers.cross_encoder.losses import BinaryCrossEntropyLoss
from sentence_transformers.cross_encoder.modules import LogitScore, Transformer
from sentence_transformers.cross_encoder.trainer import CrossEncoderTrainer
from sentence_transformers.cross_encoder.training_args import CrossEncoderTrainingArguments
# 1. 从模块构建模型
transformer = Transformer(
"Qwen/Qwen3.5-0.8B",
transformer_task="any-to-any",
model_kwargs={"torch_dtype": "bfloat16", "device_map": "auto", "attn_implementation": "flash_attention_2"},
processing_kwargs={"chat_template": {"add_generation_prompt": True}},
)
# 扩展聊天模板以支持"query"和"document"角色
transformer.processor.chat_template = transformer.processor.chat_template.replace(
'message.role == "user"', 'message.role in ["user", "query", "document"]'
)
# LogitScore: score = log(P("1")) - log(P("0"))
score_head = LogitScore(
true_token_id=transformer.tokenizer.convert_tokens_to_ids("1"),
false_token_id=transformer.tokenizer.convert_tokens_to_ids("0"),
)
model = CrossEncoder(
modules=[transformer, score_head],
num_labels=1,
prompts={
"image_to_text": "Given the image, judge whether the text matches it. Respond with 1 if they match, 0 if they don't.",
"text_to_image": "Given the text, judge whether the image matches it. Respond with 1 if they match, 0 if they don't.",
},
)
# 2. 定义损失
loss = BinaryCrossEntropyLoss(model)
# 3. 多