SWIFT 采样管线实战:test-time compute 采样、PRM/ORM 奖励过滤与 OpenAI API 大模型数据蒸馏
【免费下载链接】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(ms-swift)的 Sampling 能力可理解为 "test-time compute" 的工程化落地:对同一批数据用模型多次采样得到多个候选答案,再借助结果奖励模型(ORM)或过程奖励模型(PRM)对候选质量打分过滤,从而筛选出高质量正样本(并附带负样本),为强化微调(RFT)和 GRPO 等 RL 流程提供高质量训练数据。本文基于仓库中的 采样文档 与 采样管线源码,完整覆盖基础采样、奖励模型过滤、自定义 PRM/ORM、OOM 的两阶段规避方案,以及通过 OpenAI 兼容 API 调用大模型进行数据蒸馏的全部用法,并深入源码解释每个参数背后的真实行为。
能力概览与基础用法
采样能力的最小可用命令如下:
swift sample --model LLM-Research/Meta-Llama-3.1-8B-Instruct --sampler_engine transformers --num_return_sequences 5 --dataset AI-ModelScope/alpaca-gpt4-data-zh#5执行后,会在当前目录的sample_output目录下生成一个以时间戳为文件名的jsonl文件。以上命令处理 5 条数据、每条采样 5 次,因此文件应包含 25 行,每行都是一条完整的messages格式数据。
从 采样主管线 可以看到输出文件的完整生成机制:管线先把结果写入output_file.tmp临时文件,每处理完一个 batch 就flush并复制为.resume文件、同时把当前 batch 序号写入ckpt_state.json,全部结束后才把临时文件移动为最终的output_file。这意味着:
- 默认文件名格式为
YYYY-MM-DD-HH-MM-SS.jsonl(由 SamplingArguments.post_init在output_file为 None 时自动用时间戳生成); --output_file只允许文件名、不允许包含目录分隔符,否则直接抛出ValueError;- 若最终文件已存在且未指定
--override_exist_file,管线会直接返回,避免覆盖历史产物; --resume true时会从.resume文件和ckpt_state.json恢复中断的 batch,跳过已完成的采样批次。
采样参数全集可参考 命令行参数文档,下文结合 参数定义源码 给出逐项说明。
环境准备
通过 pip 安装:
pip install ms-swift[llm] -U或从源码安装:
git clone https://github.com/modelscope/ms-swift.git cd ms-swift pip install -e '.[llm]'采样入口为 swift/cli/sample.py,其中会先调用try_init_ray()再执行sampling_main()。采样器(VanillaSampler/DistillSampler)上的方法通过@RayHelper.worker/@RayHelper.function标注了sampler/prm/orm分组,从源码结构看,SWIFT 支持借助 Ray 把采样、PRM 打分、ORM 打分分发到不同进程/节点上执行,这也与 Ray 相关文档 所述的分布式能力一脉相承。
采样参数详解
以下参数表依据 SamplingArguments 的 dataclass 定义整理:
| 参数 | 默认值 | 说明 |
|---|---|---|
--model | 必填 | 待采样模型的 ID 或本地路径;sampler_engine=client时为远端 API 的模型名 |
--sampler_type | sample | 采样类型:sample(本地/服务端采样)或distill(OpenAI API 蒸馏) |
--sampler_engine | transformers | 采样推理引擎:transformers/vllm/lmdeploy/client/no。no表示不加载采样模型、仅做过滤(两阶段方案的第二阶段) |
--num_return_sequences | 64 | 每条数据采样的原始序列数 |
--num_sampling_batch_size | 1 | 每个采样批次的输入条数 |
--num_sampling_batches | None | 总采样批次数;不指定时处理整个数据集 |
--n_best_to_keep | 5 | 奖励打分后每条数据保留的最优候选数 |
--output_dir | sample_output | 输出目录 |
--output_file | 时间戳 | 输出文件名(仅.jsonl,不含目录) |
--resume | False | 断点续采 |
--override_exist_file | False | 目标文件已存在时是否覆盖 |
--prm_model | None | PRM:模型 ID(用 transformers 引擎加载)或插件中的 PRM key |
--orm_model | None | ORM:插件 key(如math)或奖励模型 ID |
--temperature | 1.0 | 采样温度 |
--prm_threshold | 0.0 | PRM 分数阈值,低于该值的候选被过滤 |
--easy_query_threshold | None | 单条 query 的正确采样比例超过该阈值则整条丢弃,过滤过简单的题目 |
--engine_kwargs | None | 传给采样引擎的 JSON 参数字符串,如'{"base_url":"..."}' |
--data_range | [] | 数据分片[shard_index, num_shards],如[1, 3]表示 3 片中取第 2 片(0 起始) |
关于engine_kwargs,vanilla_sampler.py 中有一段值得注意的细节:engine_kwargs里若残留torch_dtype会被弹出并忽略,提示改用全局--torch_dtype,以避免引擎构造时出现重复关键字参数报错。
使用 PRM 与 ORM 过滤采样结果
采样的核心价值在于对过程和结果的监督。在采样命令上追加--prm_model与--orm_model即可启用打分过滤:
swift sample --model LLM-Research/Meta-Llama-3.1-8B-Instruct --sampler_engine lmdeploy --num_return_sequences 5 --n_best_to_keep 2 --dataset tastelikefeet/competition_math#5 --prm_model AI-ModelScope/GRM-llama3.2-3B-rewardmodel-ft --orm_model math执行后sample_output目录中同样生成时间戳命名的jsonl文件,但其中最多包含 10 行(5 条数据 × 每条保留 2 条)。之所以是 "at most",是因为 ORM 验证可能失败:未通过掩码校验的候选会被直接丢弃。此外,加入 PRM/ORM 后文件格式会多出一个rejected_response键,存放该条数据中 PRM/ORM 得分最低的负样本,可直接用于 DPO 等偏好训练。
源码中的打分与筛选逻辑(见 VanillaSampler.do_sample):
- 对每条数据,将
num_return_sequences个采样结果外加 ground truth 一起送入 ORM 和 PRM 打分; - 综合得分
score = prm_score + orm_score * 10(ORM 权重更大); - 得分先经 get_reward 做 min-max 归一化,PRM 分数还需同时满足
prm_score > prm_threshold(ORM 侧阈值为 0,即必须大于 0 才算通过); - 按综合得分排序后取前
n_best_to_keep个通过掩码的候选作为正样本,得分最低者作为负样本写入rejected_response; - 若启用
easy_query_threshold,当一条 query 的正确候选数达到num_return_sequences * easy_query_threshold时,认为题目过于简单,整条 query 被丢弃; - 所有候选都未通过掩码时,该条数据整体跳过——这正是输出 "at most 10 行" 的原因。
内置的 ORM/PRM 注册表分别位于 swift/rewards/orm.py 与 swift/rewards/prm.py。可用的 ORM key 包括math(基于 sympy 的数学答案等价性校验)、accuracy(math_verify 校验)、format(<thinking>…</thinking><answer>…</answer>格式校验)、react_format、cosine、repetition、soft_overlong、toolbench;PRM key 包括qwen_max(调用 qwen-max 作为裁判)和client(可配置 base_url/模型的 API 裁判)。若传入的--prm_model/--orm_model不在注册表中,Sampler._prepare_prm/_prepare_orm 会把它当作模型 ID,用TransformersEngine加载一个真正的奖励模型。
自定义 PRM 或 ORM
PRM/ORM 均支持插件式扩展:按现有代码新增一个实现类并注册到对应字典,即可在命令行引用。自定义 PRM 示例:
class CustomPRM: # 构造函数必须无参 def __init__(self): # 在这里初始化 pass def __call__(self, infer_requests: List[InferRequest], ground_truths: List[str], **kwargs) -> List[Union[float, List[float]]]: ... prms = {'custom': CustomPRM}注册完成后,命令行使用--prm_model custom即可。ORM 的自定义方式类似:实现继承自 ORM 的类,__call__接收completions/infer_requests、solution/ground_truths等参数并返回 float 奖励列表(若签名中没有ground_truths参数,get_reward 会通过inspect.signature自动判断是否传入 ground truth)。
内存控制:两阶段采样避免 OOM
若采样模型与 PRM 奖励模型同时驻留显存,极易触发 OOM。SWIFT 通过--cache_files参数支持把流程拆成两阶段:
- 阶段 1:只指定
--model和--sampler_engine,不指定--orm_model/--prm_model,仅执行采样并把全部结果落盘; - 阶段 2:指定
--sampler_engine no、--orm_model和--prm_model,并用--cache_files指向阶段 1 的输出文件,只对缓存结果做 RM 过滤,不再重新采样。
阶段 2 的读取逻辑见 read_cache:按行解析缓存文件,以数据内容的 MD5(由 get_messages_md5 对去掉choices的行做sort_keys序列化后计算)作为 key,把每行的messages末条 assistant 回复还原为候选列表。需要注意源码中的两点提示:阶段 2 仍必须提供--dataset,因为缓存中的 MD5 id 需要与原始数据重新关联;若某个 query 的缓存候选数不足num_return_sequences(如阶段 1 有采样失败),该 query 会在阶段 2 重新走缓存合并/补采样路径。
实战示例:采样驱动的 RFT
采样能力是 RFT(Reinforcement Fine-Tuning)的关键环节,仓库提供了完整示例脚本 examples/train/rft/rft.py,演示"采样 → 奖励过滤 → 用筛选出的数据强化训练"的闭环。仓库中另有 examples/sampler/sample/ 与 examples/sampler/distill/ 两个目录,分别存放普通采样与蒸馏采样的 shell 脚本和配置文件,可作为命令组装的参考。
需要注意:该脚本的实际效果与模型、数据和 RM 的质量强相关,官方仅将其作为示例提供,建议读者结合自身场景修改脚本并自行训练 RM 与生成模型。
通过 OpenAI API 从大模型蒸馏数据
sample还支持以 OpenAI 兼容 API 调用远端大模型批量蒸馏数据(--sampler_type distill):
OPENAI_API_KEY="your_api_key" \ swift sample \ --sampler_type distill \ --sampler_engine client \ --model deepseek-r1 \ --stream true \ --dataset tastelikefeet/competition_math#5 \ --num_return_sequences 1 \ --temperature 0.6 \ --top_p 0.95 \ --engine_kwargs '{"base_url":"https://dashscope.aliyuncs.com/compatible-mode/v1"}'其中base_url和model分别表示 API 端点与远端模型名,stream控制请求是否流式。源码层面:sampler_engine=client时 SamplingArguments._init_model_info 会跳过本地模型/模板的加载(model_info/model_meta置 None),因此该模式不占本地显存;DistillSampler 内部使用OpenAIEngine封装openaiSDK 的chat.completions调用,api_key缺省时读取环境变量OPENAI_API_KEY,base_url缺省指向阿里云百炼的兼容端点。
一个重要的输出格式约定:对于 DeepSeek-R1 等带推理内容的模型,流式分片会先收集reasoning_content再收集content,最终拼接为固定格式:
<thinking>{reasoning_content}</thinking> <answer>{content}</answer>该格式与内置 ORMFormat的校验正则^<thinking>.*?</thinking>\s*<answer>.*?</answer>(见 orm.py)严格对应——也就是说,蒸馏得到的数据可以直接配合--orm_model format或math/accuracy做格式与答案校验,形成"API 蒸馏 + 本地过滤"的完整数据生产链路。
总结
- SWIFT 的
swift sample提供从基础多次采样、奖励模型过滤到 API 蒸馏的一站式数据生产管线,输出为可直接用于 SFT/DPO/RFT 的messages格式 jsonl; - 过滤逻辑为
score = PRM + ORM*10的加权排序,配合n_best_to_keep、prm_threshold、easy_query_threshold精细控制正负样本产出,负样本以rejected_response键随正样本一并落盘; - 两阶段
--cache_files方案与data_range分片解决了奖励模型 OOM 与多机并行两类工程问题; - 所有行为均可在 swift/arguments/sampling_args.py、swift/pipelines/sampling/ 与 swift/rewards/ 中逐行核对,便于按项目需求扩展自定义 PRM/ORM 或调整打分策略。
【免费下载链接】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),仅供参考