如何使用 SWIFT 训练 Reranker 模型?数据集格式与 pointwise/listwise 损失说明
【免费下载链接】swiftUse PEFT or Full-parameter to CPT/SFT/DPO/GRPO 600+ LLMs (Qwen3.6, DeepSeek-V4, GLM-5.1, InternLM3, Llama4, ...) and 300+ MLLMs (Qwen3-VL, Qwen3-Omni, InternVL3.5, Ovis2.5, GLM4.5v, Gemma4, Llava, Phi4, ...) (AAAI 2025).项目地址: https://gitcode.com/GitHub_Trending/swift1/swift
SWIFT 已经支持 Reranker 模型的训练,覆盖两类架构:分类式 Reranker(在预训练模型上加分类头,输出单个相关性分数,如iic/gte-reranker-modernbert-base)和生成式 Reranker(基于 CausalLM 架构,通过对比最后位置特定 token(如 "yes"/"no")的 logits 做判断,如 Qwen3-Reranker 0.6B/4B/8B 和 Qwen3-VL-Reranker 2B/8B)。本文以完成一次 Reranker 微调为主线:先准备符合 SWIFT 要求的数据集,再选择pointwise_reranker或listwise_reranker损失,最后运行仓库提供的脚手架脚本并用推理脚本验证打分效果。
选择模型时先确定任务类型(task_type)
训练命令中需要显式指定--task_type,它与模型架构对应:
- 分类式模型(如
iic/gte-reranker-modernbert-base):--task_type reranker; - 生成式模型(如
Qwen/Qwen3-Reranker-4B、Qwen/Qwen3-VL-Reranker-8B):--task_type generative_reranker。
两者的区别在于损失计算前的打分方式:分类式输出单个相关性分数;生成式输出特定 token 的概率。仓库中对应的脚手架脚本见 train_reranker.sh(分类式)和 train_generative_reranker.sh(生成式)。
准备数据集:positive_messages 与 negative_messages
SWIFT 的 Reranker 数据集为 JSON Lines 格式。纯文本(LLM)场景下,每行形如:
{"messages": [{"role": "user", "content": "query"}], "positive_messages": [[{"role": "assistant", "content": "relevant_doc1"}],[{"role": "assistant", "content": "relevant_doc2"}]], "negative_messages": [[{"role": "assistant", "content": "irrelevant_doc1"}],[{"role": "assistant", "content": "irrelevant_doc2"}]]}多模态(MLLM)场景下,messages内容里用<image>标记图片位置,并额外提供images、positive_images、negative_images字段(文档示例,其中路径为占位值,需替换为你本地实际图片路径):
{"messages": [{"role": "user", "content": "<image>query"}], "images": ["/some/images.jpg"], "positive_messages": [[{"role": "assistant", "content": "<image>relevant_doc1"}]], "positive_images": [["/some/positive_images.jpg"]], "negative_messages": [[{"role": "assistant", "content": "<image><image>irrelevant_doc1"}], [{"role": "assistant", "content": "<image>irrelevant_doc2"}]], "negative_images": [["/some/negative_images1.jpg", "/some/negative_images2.jpg"], ["/some/negative_images3.jpg"]]}字段含义:
messages:查询文本;positive_messages:与查询相关的正例文档列表,支持多个正例;negative_messages:与查询不相关的负例文档列表,支持多个负例。
采样行为由两个环境变量控制:
MAX_POSITIVE_SAMPLES:每个 query 的最大正例数量(默认 1);MAX_NEGATIVE_SAMPLES:每个 query 的最大负例数量(默认 7)。
默认情况下,每条数据会取出MAX_POSITIVE_SAMPLES条正样本和MAX_NEGATIVE_SAMPLES条负样本,每条正样本会和这些负样本组成一个 group,因此每条数据会扩展成MAX_POSITIVE_SAMPLES × (1 + MAX_NEGATIVE_SAMPLES)条数据。正例或负例数量不足时取全部;超过上限时随机采样。
这里有一个直接影响显存的细节:展开后的数据会放在同一个 batch 中,因此每个设备上的实际批处理大小为per_device_train_batch_size × MAX_POSITIVE_SAMPLES × (1 + MAX_NEGATIVE_SAMPLES)。调整per_device_train_batch_size时按这个公式核算,避免显存不足;如果不想使用默认 7 个负例,也可以在命令前显式设置MAX_NEGATIVE_SAMPLES,仓库的多模态脚本就是这样做的(见 train_reranker_mm.sh 开头的MAX_NEGATIVE_SAMPLES=1)。
选择损失函数:pointwise 还是 listwise
SWIFT 支持两种 Reranker 损失,通过--loss_type指定,实现见 swift/loss/mapping.py。
Pointwise:二分类,判断单个 pair 是否相关
--loss_type pointwise_reranker将排序问题转化为二分类问题,独立处理每个 query-document 对,损失为二分类交叉熵。文档描述其适用场景为简单高效、适合大规模数据训练;优点是训练简单,缺点忽略了文档间的相对关系。
生成式 Reranker 使用 pointwise 损失时,两个 token 由环境变量配置:
GENERATIVE_RERANKER_POSITIVE_TOKEN:正例 token(默认 "yes");GENERATIVE_RERANKER_NEGATIVE_TOKEN:负例 token(默认 "no")。
Listwise:多分类,从候选组中识别正例
--loss_type listwise_reranker将排序问题转化为多分类问题:对每个 query 的候选文档组(1 个正例 + n 个负例)做多分类,识别正例文档,损失为多分类交叉熵。文档认为它学习文档间的相对排序关系,更符合信息检索的实际需求。相关环境变量:
LISTWISE_RERANKER_TEMPERATURE:softmax 温度参数(默认 1.0);LISTWISE_RERANKER_MIN_GROUP_SIZE:最小组大小,组内文档数量小于该值时不计算损失(默认 2)。
四种组合(分类式/生成式 × pointwise/listwise)在仓库中都有现成脚本:
| 组合 | 脚本 |
|---|---|
| Pointwise 分类式 | train_reranker.sh |
| Pointwise 生成式 | train_generative_reranker.sh |
| Listwise 分类式 | train_reranker_listwise.sh |
| Listwise 生成式 | train_generative_reranker_listwise.sh |
运行训练:两条主路径
路径一:分类式模型 pointwise 训练(1 卡,约 5G 显存)
以gte-reranker-modernbert-base全参数微调为例,脚本 train_reranker.sh 可直接执行:
CUDA_VISIBLE_DEVICES=0 \ swift sft \ --model iic/gte-reranker-modernbert-base \ --task_type reranker \ --loss_type pointwise_reranker \ --tuner_type full \ --dataset MTEB/scidocs-reranking \ --load_from_cache_file true \ --split_dataset_ratio 0.05 \ --eval_strategy steps \ --output_dir output \ --eval_steps 100 \ --num_train_epochs 1 \ --save_steps 200 \ --per_device_train_batch_size 64 \ --per_device_eval_batch_size 64 \ --gradient_accumulation_steps 1 \ --dataset_num_proc 8 \ --learning_rate 6e-6 \ --label_names labels \ --dataloader_drop_last true其中--dataset MTEB/scidocs-reranking是脚手架自带的示例数据集,换成自己的数据时按上节格式准备;--split_dataset_ratio 0.05切出 5% 做验证集,配合--eval_strategy steps和--eval_steps 100按步评估。
改用 listwise 损失时,只需将--loss_type换成listwise_reranker(见 train_reranker_listwise.sh),并按需通过LISTWISE_RERANKER_TEMPERATURE、LISTWISE_RERANKER_MIN_GROUP_SIZE环境变量调参。
路径二:Qwen3-Reranker 生成式训练(4 卡,约 47G 显存)
脚本 train_generative_reranker.sh 展示 4 卡全参数 pointwise 训练:
CUDA_VISIBLE_DEVICES=0,1,2,3 \ NPROC_PER_NODE=4 \ swift sft \ --model Qwen/Qwen3-Reranker-4B \ --task_type generative_reranker \ --loss_type pointwise_reranker \ --tuner_type full \ --dataset MTEB/scidocs-reranking \ --load_from_cache_file true \ --split_dataset_ratio 0.05 \ --eval_strategy steps \ --output_dir output \ --eval_steps 100 \ --num_train_epochs 1 \ --save_steps 200 \ --per_device_train_batch_size 2 \ --per_device_eval_batch_size 2 \ --gradient_accumulation_steps 8 \ --dataset_num_proc 8 \ --learning_rate 6e-6 \ --label_names labels \ --deepspeed zero2 \ --dataloader_drop_last true注意 4B 模型下per_device_train_batch_size只有 2,并叠加 deepspeed zero2,且MAX_NEGATIVE_SAMPLES默认 7,实际每卡批大小需按前文的展开公式核算,显存紧张时应下调per_device_train_batch_size或MAX_NEGATIVE_SAMPLES。
可选分支:卡数或显存不够时,仓库提供了 LoRA 版本 qwen3_reranker.sh(2 卡约 20GiB,--tuner_type lora --lora_rank 8 --lora_alpha 32 --learning_rate 5e-5 --target_modules all-linear,并使用--attn_impl flash_attn --padding_free true);多模态训练可参考 qwen3_vl_reranker.sh(数据集swift/TextCaps:rerank)和 train_reranker_mm.sh。
验证:用推理脚本查看打分
训练结束后,用仓库的推理脚本检查模型是否输出相关性分数。全参数训练参考 demo_reranker.py:
import torch from swift.infer_engine import InferRequest, TransformersEngine engine = TransformersEngine( 'Qwen/Qwen3-Reranker-4B', task_type='generative_reranker', torch_dtype=torch.float16, attn_impl='flash_attention_2') infer_request = InferRequest( messages=[{ 'role': 'system', 'content': 'Given a web search query, retrieve relevant passages that answer the query' }, { 'role': 'user', 'content': 'What is the capital of China?' }, { 'role': 'assistant', 'content': 'The capital of China is Beijing.' }]) response = engine.infer([infer_request])[0] print(f'scores: {response.choices[0].message.content}')LoRA 微调后的 checkpoint 加载方式见 qwen3/infer.py,通过TransformersEngine的adapters参数指向训练输出目录,文档示例为adapters=['output/vx-xxx/checkpoint-xxx'],其中的output/vx-xxx/checkpoint-xxx需要替换为本次训练实际生成的 checkpoint 路径:
engine = TransformersEngine( 'Qwen/Qwen3-Reranker-4B', task_type='generative_reranker', attn_impl='flash_attention_2', adapters=['output/vx-xxx/checkpoint-xxx']) # 替换为实际训练输出的 checkpoint 路径该脚本会对同一 query 下的相关文档和不相关文档分别请求,打印各自的scores,可据此判断微调后模型是否对正例给出更高分。
进阶:Qwen3-Reranker 自定义 Instruction
Qwen3-Reranker 默认 Instruction 为Given a web search query, retrieve relevant passages that answer the query。如果任务场景不是通用网页检索,可以在数据中覆盖:positive_messages/negative_messages内提供的system优先于主messages的system,两者都未提供时使用默认 Instruction。这一"就近覆盖"规则在多份正负例的 dataset 中同样适用。
限制与注意点
- 文档支持列表中目前列出的 Reranker 模型为 modernbert reranker、qwen3-reranker(0.6B/4B/8B)和 qwen3-vl-reranker(2B/8B),其他模型需自行核对是否可注册;
- 展开后数据同 batch 的机制意味着默认
MAX_NEGATIVE_SAMPLES=7时批大小会被放大 8 倍(以MAX_POSITIVE_SAMPLES=1计),OOM 时优先调小这两个环境变量或per_device_train_batch_size; - listwise 损失下,组内文档数小于
LISTWISE_RERANKER_MIN_GROUP_SIZE(默认 2)时该组不计算损失,负例极少的数据会显著损失有效样本; - 完整说明以 Reranker 训练文档 为准,本文的命令均摘自仓库
examples/train/reranker/下的脚手架脚本。
【免费下载链接】swiftUse PEFT or Full-parameter to CPT/SFT/DPO/GRPO 600+ LLMs (Qwen3.6, DeepSeek-V4, GLM-5.1, InternLM3, Llama4, ...) and 300+ MLLMs (Qwen3-VL, Qwen3-Omni, InternVL3.5, Ovis2.5, GLM4.5v, Gemma4, Llava, Phi4, ...) (AAAI 2025).项目地址: https://gitcode.com/GitHub_Trending/swift1/swift
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考