尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

ms-swift GRPO 多轮训练指南:用 MultiTurnScheduler 定制多轮 Rollout 流程

ms-swift GRPO 多轮训练指南:用 MultiTurnScheduler 定制多轮 Rollout 流程 ms-swift GRPO 多轮训练指南用 MultiTurnScheduler 定制多轮 Rollout 流程【免费下载链接】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在强化学习训练中模型的采样往往不是一次生成就能完成的工具调用、环境交互、纠错反思等场景都要求模型与环境进行多轮往返并根据环境反馈持续推理。ms-swift 通过MultiTurnScheduler抽象基类和swift rollout服务命令为 GRPO以及共享同一套基础设施的 GKD提供了可扩展的多轮训练框架。读完本文你将能够理解多轮 rollout 的默认调度逻辑掌握multi_turn_scheduler、max_turns、vllm_use_async_engine等关键参数的配置方式并学会用 hook 或完全重写run方法实现自己的多轮交互逻辑包括损失掩码、奖励信息回传与训推一致性修正。多轮训练的整体流程多轮训练的核心差异在于一条训练轨迹不再由单次采样构成而是由「模型生成 → 环境反馈 → 模型再生成」的多轮循环构成循环可能包含环境交互、工具调用等步骤直到满足终止条件为止。框架把这一循环抽象为多轮规划器Scheduler由它统一负责对话状态的推进与终止判断。规划器承担两大核心功能终止条件判断通过check_finished方法判断当前轮次推理是否应该结束推理请求构造通过step方法构建下一轮推理的请求对象。MultiTurnScheduler 抽象基类及其核心方法MultiTurnScheduler定义在 swift/rollout/multi_turn.py是RolloutScheduler的子类。源码中给出了两类定制路径的说明完全定制直接重写run()方法获得对 rollout 过程的全部控制权需自行处理轮次管理与终止逻辑部分定制实现必需的step()方法并可选重写check_finished()复用基类run()提供的轮次管理基础设施。基类的核心接口如下与文档一致源码中 hook 为async方法以便直接await异步 gym 环境class MultiTurnScheduler(ABC): def __init__(self, max_turns: Optional[int] None, *args, **kwargs): self.max_turns max_turns def on_trajectory_start(self, requests: List[RolloutInferRequest]) - None: 在首轮推理前调用用于初始化轨迹级别状态。 可在此方法中直接修改 requests如注入环境初始 observation。 默认实现为空no-op。 pass def on_turn_end(self, infer_request: RolloutInferRequest, response_choice: ChatCompletionResponseChoice, current_turn: int) - Dict[str, Any]: 在 assistant 消息追加后、check_finished 前调用。 用于推进环境状态如 env.step并返回每轮元数据。 返回值可选包含 - done (bool): 若存在将覆盖 check_finished 的结果 - rollout_infos (dict): 合并到轨迹累积的额外信息中 默认返回空字典no-op。 return {} def step(self, infer_request: RolloutInferRequest, response_choice: ChatCompletionResponseChoice, current_turn: int) - Dict: 处理对话轮次之间的转换返回下一轮推理请求等结果 - infer_request (必需): 下一轮的推理请求对象 - response_token_ids (可选): 每个 rollout 轮次的响应 token IDs - response_loss_mask (可选): 每个 rollout 轮次响应的损失掩码 - rollout_logprobs (可选): 每个 rollout 轮次的响应对应的 logps - rollout_infos (可选): 额外信息数据 raise NotImplementedError def check_finished(self, infer_request: RolloutInferRequest, response_choice: ChatCompletionResponseChoice, current_turn: int) - bool: 默认终止逻辑 1. 当响应达到长度限制时 (finish_reason length) 2. 当对话达到最大轮数时 (如果设置了 max_turns) if response_choice.finish_reason length: return True if self.max_turns and current_turn self.max_turns: return True return Falsestep/check_finished等方法的入参说明infer_request当前的推理请求RolloutInferRequest包含messages、data_dict数据集中的原始列等字段response_choice当前轮次的推理结果ChatCompletionResponseChoice包含生成的message.content、finish_reason、token_ids、logprobs等current_turn当前推理轮次从 1 开始。一个典型的入参形态来自文档中的入参示例此处摘录结构# infer_request RolloutInferRequest( messages[ {role: system, content: ...think about the reasoning process... answer ... /answer...}, {role: user, content: What is the value of $\\sqrt{36 \\times \\sqrt{16}}$?}, {role: assistant, content: To find the value of ... \\boxed{12}} ], images[], audios[], videos[], toolsNone, objects{}, data_dict{problem: ..., solution: ...} ) # response_choice ChatCompletionResponseChoice( index0, messageChatMessage(roleassistant, content..., tool_callsNone), finish_reasonstop, logprobsNone, ) # response_choice.messages 会在多轮推理结束时被复制默认的check_finished会在以下两种情况停止推理模型回复被截断即超出了max_completion_lengthfinish_reason length模型推理轮数超出了max_turns限制。完整的默认多轮 rollout 逻辑在基类的run方法中见 swift/rollout/multi_turn.py。从源码实现看默认run的循环为每轮先调用推理引擎生成把 completion 追加进infer_request.messages随后触发on_turn_endhook、执行check_finishedon_turn_end返回的done会覆盖该结果并额外做一次max_turns的兜底检查防止用户忘记判断轮数上限源码 L409-L411若未终止则调用step得到下一轮请求并把每轮的response_token_ids、response_loss_mask、rollout_logprobs按轮累积最终打包进RolloutOutput返回。值得注意的实现细节hook 双模式通用on_trajectory_start/on_turn_end是通用 hook同时被 server mode 的run()与 colocate mode 的run_multi_turn()见 swift/rollout/agent_loop.py调用因此基于 hook 的自定义逻辑天然兼容两种部署模式token 精确性get_response_token_data会把模板的 response prefix如 assistant 起始标记以掩码 0 的 ID 形式补入训练 completion且断言response_loss_mask与response_token_ids等长、取值仅为 0/1logprobs 完整性校验若累积的rollout_logprobs数量与loss_mask1的 token 数不一致框架会主动清空 logprobs从而自动禁用 rollout 重要性采样修正源码 L434-L453保证不会用错位的数据做 off-policy 校正。框架内置了四个可直接使用的规划器注册在 swift/rollout/multi_turn.py 的multi_turns表中注册名实现类典型用途math_tip_trickMathTipsScheduler答案错误时注入一次提示让模型重新检查推理thinking_tips_schedulerThinkingModelTipsScheduler思考类模型多轮推理每轮历史只保留最后一轮思考gym_schedulerGYMScheduler标准 gym 环境reset/step/环境直接给奖励openenv_schedulerOpenEnvSchedulerOpenEnv 同步 WebSocket 环境动作解析/观测格式化可覆写设置多轮训练参数多轮训练采用 server mode先用swift rollout启动带规划器的 rollout 服务再由swift rlhf通过 vLLM server 连接它。rollout 端通过multi_turn_scheduler参数指定规划器名称通过max_turns限制最大轮数swift rollout \ --model Qwen/Qwen3-1.7B \ --vllm_use_async_engine true \ --multi_turn_scheduler thinking_tips_scheduler \ --vllm_max_model_len 32768 \ --vllm_gpu_memory_utilization 0.8 \ --max_turns 3参数与校验逻辑可从源码得到确认swift/arguments/deploy_args.pymulti_turn_scheduler默认None指定multi_turns注册表中的规划器名称max_turns默认None即不限制最大轮数vllm_use_async_engine若已设置multi_turn_scheduler而未显式指定该参数框架会默认置为True若两者矛盾设置了调度器但 async engine 为 False会直接抛出ValueError。服务端的调度器实例化发生在 swift/pipelines/infer/rollout.py 的get_rollout_engine_type中调度器名不在multi_turns表里会报not found错误实例化时传入infer_engine、max_turns并按构造函数签名自动注入tokenizergym 场景还会带上use_gym_env/gym_env参数。通过 external_plugins 注册本地规划器内置表之外的规划器可以通过external_plugins参数把本地 Python 文件注册进 ms-swift在插件文件中定义规划器子类并写入multi_turns注册表然后在 rollout/训练命令中加--external_plugins 插件文件。参考实现是 examples/train/grpo/plugin/plugin.py其中包含ToolCallScheduler工具调用场景演示response_loss_mask用法等规划器。一个完整的多轮训练脚本见 examples/train/grpo/external/vllm_multi_turn.sh。该示例演示的是一条多轮轨迹拆分为多条数据的场景对应thinking_tips_scheduler先用swift rollout启动服务脚本前半段已注释再以 GRPO 方式训练关键参数如下CUDA_VISIBLE_DEVICES1,2 \ NPROC_PER_NODE2 \ swift rlhf \ --rlhf_type grpo \ --model Qwen/Qwen3-1.7B \ --tuner_type full \ --external_plugins examples/train/grpo/plugin/plugin.py \ --reward_funcs thinking_tips \ --loss_scale last_round \ --use_vllm true \ --vllm_mode server \ --vllm_server_host 127.0.0.1 \ --vllm_server_port 8000 \ --vllm_server_pass_dataset true \ --dataset AI-MO/NuminaMath-TIR#10000 \ --max_completion_length 8192 \ --num_generations 8 \ --num_iterations 1 \ --importance_sampling_level sequence \ --deepspeed zero2 \ ...脚本注释明确了两点实战约束同一轨迹拆出的多条数据奖励函数需要给出相同奖励示例用每条轨迹最后一轮数据计算 accuracy 奖励--vllm_server_pass_dataset true用于把数据集额外列传入规划器的data_dict详见下文。AsyncEngine 与多轮采样效率多轮 rollout 中不同轨迹的轮数往往不同有的 1 轮就终止有的跑满max_turns。若用同步批量推理必须等最慢轨迹完成才能发下一批请求产生大量计算气泡。AsyncEngine 通过异步并发调度让各轨迹独立推进减少多轮推理过程中的计算气泡。使用方式就是在rollout命令中设置vllm_use_async_engine默认使用 async engine。注意async engine 仅在 server mode 下可用这也是MultiTurnScheduler.run方法重写仅在swift rolloutvllm_use_async_engineTrue时生效的原因。GYM 环境训练如果你的多轮任务可以建模为标准的 gym environmentreset/step/ 环境直接给奖励推荐直接复用框架内置的gym_scheduler只需实现一个Env子类定义在 swift/rollout/gym_env.py来描述任务。GYMScheduler基于通用 hook 协议实现源码 swift/rollout/multi_turn.py无需重载run方法on_trajectory_start为每个请求reset一个环境实例把初始 observation及可选 system message注入首轮 user 消息on_turn_end调用env.step(messages)推进环境累积total_reward/step_rewards返回{done: bool, rollout_infos: {total_reward, step_rewards, gym_done}}step把on_turn_end暂存的下一 observation 作为新的 user 消息追加进对话。这种设计使得GYMScheduler同时适用于 server moderun()和 colocate moderun_multi_turn()用户只需实现 Env 接口。环境名可来自命令行--gym_env也可来自数据集每行的env_config列data_dict[env_config]。完整接口与自定义 env 的步骤可参考同目录的 gym_env.md 文档。若环境是 OpenEnv 的同步 WebSocket 服务则用openenv_scheduler它额外提供parse_actionLLM 文本→动作 dict与format_observation观测 dict→字符串两个可覆写方法。高级设置自定义多轮交互逻辑默认逻辑把一条轨迹的整体消息历史用来计算多轮 rollout 的损失这里隐含一个假设多轮交互过程中模型的历史信息没有受到改变。但在一些多轮场景中需要在 rollout 过程中动态修改模型的历史信息比如压缩历史信息、按轮裁剪此时需要把每轮的 rollout 单独作为一条轨迹进行训练。方式一使用 hook利用on_trajectory_start/on_turn_end两个 hook 即可表达大部分环境生命周期逻辑同时适用于 server mode 和 colocate mode无需重载run方法class CustomScheduler(MultiTurnScheduler): def on_trajectory_start(self, requests): # 首轮推理前初始化如环境 reset、注入初始状态 for req in requests: req.messages [system_msg, user_msg(initial_observation)] def on_turn_end(self, req, response_choice, current_turn): # 每轮推理后推进状态返回 done 和 rollout_infos next_obs, reward, done self.advance_env(req.messages) return { done: done, rollout_infos: {reward: reward, ...} }框架还会做一次max_turns兜底run中强制should_stop or current_turn max_turns所以即使自定义check_finished忘记判断轮数也不会无限循环。方式二重载 run 方法完全自定义一个常见场景是思考类模型实际推理中模型通常只保留最后一轮的思考内容而忽略历史回复中的思考内容。这类场景需要重写规划器的交互逻辑即重载run方法把每一轮单独作为一个RolloutOutput返回。框架内置的ThinkingModelTipsSchedulerswift/rollout/multi_turn.py演示了完整做法每轮生成后_build_messages基于模板的 thinking 处理逻辑重建该轮的历史仅保留最后一轮的 think 内容并把该轮结果 append 进rollout_outputs列表返回答案正确或已给过提示时终止。注意这种方式下相同轨迹的数据会拆分为多条数据在奖励相关处理中必须对同一轨迹的数据分配相同的 reward可通过kwargs中的trajectory_inputs获取完整轨迹数据参考 examples/train/grpo/plugin/plugin.py 中MultiTurnThinkingTips的实现。另外源码头注释提示此类拆分场景通常要配合--loss_scale last_round只训练最后一轮回复。多模态数据修改多模态多轮交互场景下可能需要在对话过程中动态增删或修改多模态数据并确保变更同步至 trainer。实现方式是借助rollout_infos在其中写入指定键即可覆盖原始数据集的多模态内容。目前已支持覆盖的键为images、audios、videos。可参考 examples/train/grpo/plugin/deepeyes/deepeyes_plugin.py 中 DeepEyes Scheduler 的实现。返回 response token ids默认流程中规划器把模型生成的文本字符串返回给 trainertrainer 再将其重新 encode 为 token id 用于训练。为避免这一步重复编码同时规避 encode/decode 不对称带来的训推偏差可以让规划器直接返回response_token_ids在response_choice对象中读取token_ids属性即本次 rollout 生成的 token 序列在step/run方法的返回值中加入response_token_idstrainer 便能直接使用这些 token id 参与训练无需重新编码。参考实现swift/rollout/multi_turn.py 中的ThinkingModelTipsSchedulerrollout_outputs中携带response_token_idsresponse_choice.token_ids。损失掩码在工具调用或环境交互返回结果时若需将返回内容并入模型响应文本建议对插入内容做掩码确保训练时不对外部生成内容计算损失。两种设置方式第一种设置 loss_scale。ms-swift 提供loss_scale参数对模型回复部分的内容做损失处理例如--loss_scale last_round可将非最后一轮的模型回复损失置零。也可实现自定义 loss_scale参考 Customization/Architecture.md 中 loss_scale 章节。注意在 GRPO 中loss_scale 只提供掩码功能不提供缩放功能。第二种设置 response_loss_mask。在step或run方法中返回response_loss_mask即可在规划器侧完全自定义逐 token 的损失掩码。前提是必须同时返回response_token_ids且两者等长、取值 0/1源码中的断言会强制校验。返回response_loss_mask时loss_scale参数失效。参考 examples/train/grpo/plugin/plugin.py 中ToolCallScheduler的实现MathTipsScheduler也演示了给注入的提示 token 打 0 掩码的完整写法原始 token 掩码 1 提示 token 掩码 0。奖励函数相关在奖励函数中获取多轮 rollout 信息的方式在on_turn_end或step/run中返回rollout_infos然后在奖励函数的kwargs中读取class Scheduler(): def on_turn_end(self, infer_request, response_choice, current_turn): ... return {done: done, rollout_infos: extra_dict} # 或者在 step 方法中 def step(self, infer_request, response_choice, current_turn): ... return {infer_request: infer_request, rollout_infos: extra_dict} class RewardFunction(): def __call__(self, completions, **kwargs): infos kwargs.get(rollout_infos, {}) ...此外swift/rlhf_trainers/grpo_trainer.py 会在日志指标中自动收集rollout_infos里的num_turns默认run方法总是写入num_turns轮数信息便于监控实际多轮分布。在 Scheduler 中获取额外的数据集信息在训练侧设置参数--vllm_server_pass_dataset可将数据集中的其他列传入多轮规划器在infer_request.data_dict中读取。内置的ThinkingModelTipsScheduler/MathTipsScheduler正是依赖data_dict[solution]来判分决定是否继续追问的。训推一致性兼容swift 支持从 vLLM 侧返回 rollout 的 logps 用于纠正训推不一致问题背景与参数细节可参考 training_inference_mismatch.md。在多轮训练中如果启用了rollout_importance_sampling_mode框架会自动收集每轮 rollout 的 log probabilities用于校正训推不一致带来的 off-policy 问题。默认行为使用默认run方法时框架会自动从response_choice.logprobs中提取 log probabilities并与response_token_ids、response_loss_mask一起传给 trainer。自定义 Scheduler 的注意事项如果step方法修改了 response如截断、添加内容需同步返回对应的rollout_logprobs。关键规则与源码中的校验逻辑一致rollout_logprobs的长度应等于response_loss_mask中值为 1 的数量即只对参与损失计算的 token 提供 logprob而不是 response 的全部 token对loss_mask0的 token如用户添加的提示、工具返回结果不需要提供 logprobs如果step未返回rollout_logprobs框架会自动从response_choice.logprobs中提取如果 logprobs 不完整数量对不上框架会清空它们并自动禁用修正避免错位数据污染训练。重写run方法的场景如果你完全重写了run方法需要手动收集并传递rollout_logprobs可按轮次累积最终放在RolloutOutput.rollout_logprobs中。具体实现可参考 swift/rollout/multi_turn.py 中的内置实现。小结ms-swift 的多轮 GRPO 训练以MultiTurnScheduler为枢纽check_finished管终止、step管轮次推进、on_trajectory_start/on_turn_end管轨迹级与环境级状态run则提供从单轨迹整体损失到逐轮拆分轨迹的完全控制权。轻量需求优先用 hook 与gym_scheduler需要动态改写历史时重写run并配合response_token_ids/response_loss_mask/rollout_logprobs三件套保证逐 token 的训练精度与训推一致性。所有相关代码集中在 swift/rollout/multi_turn.py示例脚本与插件位于 examples/train/grpo 下可作为二次开发的起点。【免费下载链接】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),仅供参考
返回列表