news 2026/9/15 12:51:39

如何使用 SWIFT 训练 Reranker 模型?数据集格式与 pointwise/listwise 损失说明

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
如何使用 SWIFT 训练 Reranker 模型?数据集格式与 pointwise/listwise 损失说明

如何使用 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_rerankerlistwise_reranker损失,最后运行仓库提供的脚手架脚本并用推理脚本验证打分效果。

选择模型时先确定任务类型(task_type)

训练命令中需要显式指定--task_type,它与模型架构对应:

  • 分类式模型(如iic/gte-reranker-modernbert-base):--task_type reranker
  • 生成式模型(如Qwen/Qwen3-Reranker-4BQwen/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>标记图片位置,并额外提供imagespositive_imagesnegative_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_TEMPERATURELISTWISE_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_sizeMAX_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,通过TransformersEngineadapters参数指向训练输出目录,文档示例为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优先于主messagessystem,两者都未提供时使用默认 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/15 12:50:43

AI生产力革命:从基础模型到行业应用的全景解析

1. AI革命&#xff1a;生产力范式的历史性跃迁当AlphaGo在2016年击败李世石时&#xff0c;很多人还认为这只是个游戏领域的突破。但今天&#xff0c;AI已经像电力一样渗透到每个行业——从程序员用GitHub Copilot自动补全代码&#xff0c;到设计师用MidJourney生成创意方案&…

作者头像 李华
网站建设 2026/9/15 12:49:54

GIS坐标系转换、空间校正与地理配准:多源数据整合全流程实操

你有没有遇到过这种情况&#xff1a;手里有一张扫描出来的老地形图、一套CAD导出的地块边界数据、再加上GPS采集的几十个样点&#xff0c;三份数据放到GIS软件里&#xff0c;怎么也叠不到一起去。要么位置隔了几百米&#xff0c;要么形状对不上&#xff0c;要么根本显示不到同一…

作者头像 李华
网站建设 2026/9/15 12:49:38

FPGA上实现net22+WRS+Flash高精度时间同步系统

简介&#xff1a;本资源是一套基于Verilog语言实现的FPGA SPI闪存驱动与测试工程&#xff0c;面向数字电路设计初学者及FPGA开发工程师&#xff0c;解决嵌入式系统中SPI Flash擦除、页写入与字节读取等核心操作的硬件逻辑实现问题。项目已通过综合与仿真验证&#xff0c;可直接…

作者头像 李华