十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

如何在 SWIFT GRPO 中实现自定义多轮 rollout 调度 MultiTurnScheduler?

如何在 SWIFT GRPO 中实现自定义多轮 rollout 调度 MultiTurnScheduler? 如何在 SWIFT GRPO 中实现自定义多轮 rollout 调度 MultiTurnScheduler【免费下载链接】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在 swiftms-swift的 GRPO 训练中如果模型采样需要与环境进行多轮交互例如工具调用一次推理无法完成整条轨迹。此时需要自定义一个多轮规划器MultiTurnScheduler把它注册到 rollout 服务中再接入swift rlhf训练。本文给出从编写 Scheduler 类、注册、启动 rollout 服务到跑通 GRPO 训练的完整路径并说明损失掩码、奖励函数信息传递等定制点。GKD 同样支持多轮训练与 GRPO 共享同一套MultiTurnScheduler基础设施。扩展点MultiTurnScheduler 的机制与可重写方法MultiTurnScheduler是抽象基类实现在 swift/rollout/multi_turn.py负责多轮对话管理主要承担两个核心功能终止条件判断通过check_finished判断当前轮次推理是否应该结束推理请求构造通过step构建下一轮推理的请求对象。默认终止逻辑check_finished的基类实现在两种情况下停止响应达到长度限制即finish_reason length超出了max_completion_length对话达到最大轮数如果设置了max_turns。step与check_finished接收三个参数参数含义infer_request当前的推理请求RolloutInferRequest含messages、data_dict等response_choice当前轮次的推理结果ChatCompletionResponseChoice含message.content、token_ids、finish_reason等current_turn当前推理轮次从 1 开始有两种自定义方式二选一部分定制实现step方法可选地重写check_finished复用基类run()的多轮管理基础设施完全定制直接重载run()方法自行控制整个多轮流程。官方内置的ThinkingModelTipsScheduler就是重载run()的示例适用于思考类模型每轮只保留最后一轮思考内容、把每轮 rollout 拆成独立轨迹这类需要动态修改历史信息的场景。step的返回值是一个字典结构如下infer_request为必需键infer_request必需下一轮的推理请求对象response_token_ids可选每个 rollout 轮次的响应 token IDs返回后 trainer 可跳过对 completion 文本的重新编码response_loss_mask可选与response_token_ids等长的损失掩码rollout_logprobs可选响应 token 的 log probabilitiesrollout_infos可选额外元数据需可序列化会累积到轨迹级别供奖励函数读取。此外还有两个通用 hook无需重载run即可注入状态逻辑on_trajectory_start首轮推理前初始化可直接修改requests如注入环境初始 observation和on_turn_endassistant 消息追加后、check_finished前调用返回{done: bool, rollout_infos: dict}其中done会覆盖check_finished的结果。框架内置的调度器注册在multi_turns表中见 swift/rollout/multi_turn.py 末尾包括math_tip_trick、gym_scheduler、openenv_scheduler、thinking_tips_scheduler。自定义调度器通过external_plugins参数把本地文件注册进 ms-swift。编写并注册自定义 Scheduler以 ToolCallScheduler 为例官方在 examples/train/grpo/plugin/plugin.py 中提供了完整的定制说明与可运行示例ToolCallScheduler一个支持 ReAct 格式计算器工具调度的规划器。定制流程分三步引自该文件中的官方注释定义 Scheduler 类实现step必需与check_finished可选或重载run方法把类加入multi_turns注册表multi_turns[my_scheduler] MyScheduler通过--external_plugins与--multi_turn_scheduler参数启用。ToolCallScheduler的核心部分如下工具执行与解析的辅助方法_calculator_tool、_extract_tool_calls、_execute_tools省略完整实现见仓库中该文件from swift.rollout.multi_turn import MultiTurnScheduler, multi_turns from swift.infer_engine.protocol import ChatCompletionResponseChoice, RolloutInferRequest from typing import Dict class ToolCallScheduler(MultiTurnScheduler): # A simple scheduler that supports tool calls by overriding the step method # Tool parsing uses the ReAct format def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) # A simple tool registry. Extend or replace with your own tools as needed. self.tools { calculator: self._calculator_tool, } def check_finished(self, infer_request: RolloutInferRequest, response_choice: ChatCompletionResponseChoice, current_turn: int) - bool: completion response_choice.message.content tool_calls self._extract_tool_calls(completion) if tool_calls is None: return True return super().check_finished(infer_request, response_choice, current_turn) def step(self, infer_request: RolloutInferRequest, response_choice: ChatCompletionResponseChoice, current_turn: int) - Dict: completion response_choice.message.content token_ids, loss_mask, rollout_logprobs self.prepare_response_continuation(response_choice) tool_calls self._extract_tool_calls(completion) tool_results self._execute_tools(tool_calls) # append tool result to the completion infer_request.messages[-1][content] (tool_results[0]) tokenizer self.tokenizer result_tokens tokenizer.encode(tool_results[0], add_special_tokensFalse) token_ids.extend(result_tokens) loss_mask.extend([0] * len(result_tokens)) return { infer_request: infer_request, response_token_ids: token_ids, response_loss_mask: loss_mask, rollout_logprobs: rollout_logprobs, rollout_infos: { tool_results: tool_results[0], num_turns: current_turn, } } multi_turns[tool_call_scheduler] ToolCallScheduler几个值得注意的实现细节check_finished在解析不到工具调用时直接终止否则回退到基类默认的轮数/长度判断step中把工具执行结果追加到 assistant 消息里作为下一轮的历史同时把结果 token 以loss_mask0拼进response_token_ids保证模型不对外部生成的工具结果计算损失返回的rollout_infos会在轨迹级别累积奖励函数可通过kwargs中的rollout_infos读取。注意官方提示示例中的 calculator 工具只支持基础四则运算可能无法解决数据集中所有数学题实际使用时应替换为自己的工具。启动 rollout 服务并接入 GRPO 训练多轮 rollout 依赖swift rollout启动的 rollout 服务训练侧通过--vllm_mode server连接。官方示例 examples/train/grpo/plugin/run_external_scheduler.sh 给出的标准流程是先跑swift rollout再跑swift rlhf。该脚本注释明确说明需要 main branch 的 ms-swift且Before running this script, please run the followingswift rolloutscript first。第一步启动 rollout 服务CUDA_VISIBLE_DEVICES0 \ swift rollout \ --model Qwen/Qwen2.5-7B-Instruct \ --vllm_use_async_engine true \ --external_plugins examples/train/grpo/plugin/plugin.py \ --multi_turn_scheduler tool_call_scheduler \ --vllm_max_model_len 8192 \ --vllm_gpu_memory_utilization 0.8 \ --max_turns 5参数说明--multi_turn_scheduler指定规划器取值为multi_turns注册表中的键--external_plugins把包含自定义 Scheduler以及可选的奖励函数的本地 Python 文件注册进 ms-swift--max_turns 5最大多轮数与默认check_finished的轮数上限配合--vllm_use_async_engine true使用 AsyncEngine 做批量数据异步多轮采样减少多轮推理过程中的计算气泡。文档注明 async engine 仅在 server mode 下可用。多轮推理的方法重载自定义run/step仅在 server modeswift rollout且vllm_use_async_engineTrue下受支持因此多轮训练默认走 server 模式。第二步运行 GRPO 训练rollout 服务默认监听 8000 端口后启动训练摘自同一示例脚本其中SYSTEM_PROMPT为脚本中定义的 ReAct 工具说明此处保留原变量引用完整内容见脚本CUDA_VISIBLE_DEVICES1,2,3 \ NPROC_PER_NODE3 \ swift rlhf \ --rlhf_type grpo \ --model Qwen/Qwen2.5-7B-Instruct \ --reward_funcs accuracy \ --tuner_type full \ --torch_dtype bfloat16 \ --use_vllm true \ --vllm_mode server \ --vllm_server_host 127.0.0.1 \ --vllm_server_port 8000 \ --dataset AI-MO/NuminaMath-TIR#1000 \ --load_from_cache_file true \ --max_completion_length 2048 \ --num_train_epochs 1 \ --per_device_train_batch_size 1 \ --learning_rate 1e-5 \ --gradient_accumulation_steps 4 \ --eval_steps 100 \ --save_steps 100 \ --save_total_limit 2 \ --logging_steps 5 \ --output_dir output \ --warmup_ratio 0.05 \ --dataloader_num_workers 4 \ --dataset_num_proc 4 \ --num_generations 4 \ --temperature 0.9 \ --system $SYSTEM_PROMPT \ --log_completions true \ --deepspeed zero3 \ --stop_words Observation: \ --report_to swanlab tensorboard要点--vllm_mode server加上--vllm_server_host/--vllm_server_port让训练连接第一步启动的 rollout 服务Scheduler 只运行在 rollout 服务进程里所以--multi_turn_scheduler和--external_plugins注册 Scheduler 的部分只出现在swift rollout命令中--system传入描述工具调用格式ReAct 的 Thought/Action/Action Input/Observation的提示模型输出才能被ToolCallScheduler的解析逻辑识别--stop_words Observation:控制生成截止点如果训练侧还需要自定义奖励函数再在swift rlhf命令中传--external_plugins和--reward_funcs参考 examples/train/grpo/external/vllm_multi_turn.sh该示例注册了thinking_tips奖励函数并用--loss_scale last_round训练最后一轮响应。结果验证与已知报错示例脚本开启了--log_completions true用于在训练日志中查看模型生成的 completions这是验证多轮交互是否按预期发生的直接手段。另外结合文档说明可以按以下依据判断配置是否正确生效损失掩码长度断言若response_loss_mask与response_token_ids长度不一致框架会抛出response_loss_mask must have the same length as response_token_ids只想返回掩码而忘记返回 token ids 时会触发You must return response_token_ids if you want to return response_loss_mask。遇到这两条报错时应检查step的返回值rollout_infos 覆盖行为在run()的默认流程中step返回的rollout_infos是每步覆盖式更新而非累积追加需要保留每步细节时应自行改为合并逻辑轨迹拆分时的奖励一致性如果通过重载run把同一条轨迹拆分为多条训练数据如 thinking tips 场景奖励相关处理必须对相同轨迹的数据分配同样的 reward否则优势计算会失真loss_scale 在 GRPO 中只提供掩码功能不提供缩放功能。可选定制点以下定制点均出自多轮训练文档 multi_turn.md按需选用返回 response token ids 省去重复编码在response_choice中读取token_ids属性并在step/run返回值中加入response_token_idstrainer 即可直接使用避免对 completion 文本再次 encode。ToolCallScheduler即采用此方式两种方式设置损失掩码一是训练侧参数--loss_scale last_round把非最后一轮的模型回复损失置零二是step/run中返回response_loss_mask。注意一旦返回response_loss_maskloss_scale参数失效奖励函数读取多轮信息在on_turn_end或step/run中返回rollout_infos奖励函数通过kwargs.get(rollout_infos, {})获取向 Scheduler 传递数据集列训练侧设置--vllm_server_pass_dataset true即可在infer_request.data_dict中获取数据集的其他列MathTipsScheduler就是从中读取solution字段判断答案正确性的多模态数据修改借助rollout_infos指定键值可覆盖原始数据集的多模态内容现已支持的键为images、audios、videos训推一致性启用rollout_importance_sampling_mode时框架会自动收集每轮 rollout 的 log probabilities。若在step中修改了 response截断、添加内容需要同步返回对应的rollout_logprobs其长度应等于response_loss_mask中值为 1 的数量loss_mask0的 token如工具返回结果不需要提供 logprobs。若完全重写run方法则需手动收集并传递rollout_logprobsGYM 环境替代路径如果多轮任务可以建模为标准的 gym environmentreset/step/环境直接给奖励无需自己写step直接复用内置gym_scheduler并实现一个Env子类即可接口与步骤见 gym_env.md。限制示例脚本要求 main branch 的 ms-swift自定义多轮逻辑方法重载仅在 server mode 且vllm_use_async_engineTrue时受支持示例中的 calculator 工具只支持基础四则运算_extract_tool_calls按 ReAct 格式Action:/Action Input:解析更换工具或输出格式时需同步修改这两处ToolCallScheduler的设计假设每轮至多一次工具调用脚本源码中保留了对应的断言注释。完成上述步骤后训练日志中可通过log_completions查看每条样本的多轮交互过程调度器行为不符合预期时优先核对step返回键的完整性infer_request必需、掩码与 token ids 等长以及check_finished的终止条件是否覆盖了自己的任务语义。【免费下载链接】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),仅供参考
返回列表