
torchtitan Search-R1 示例用 search 工具构建多轮检索增强 GRPO 训练流水线的完整指南【免费下载链接】torchtitanA PyTorch native platform for training generative AI models项目地址: https://gitcode.com/GitHub_Trending/to/torchtitan本篇以 torchtitan RL 实验中的 Search-R1 示例torchtitan/experiments/rl/examples/search_r1为主体讲清一个模型会调用 search 工具做多轮检索问答的 GRPO 强化学习配方从数据集流、检索服务端到端、精确匹配EM奖励函数到完整的配置注册与启动命令。读完后你能理解该示例如何完全复用 torchtitan 框架自带的多轮 rollouter 与连续批处理 generator仅靠一个示例文件夹加配置即可跑通并掌握启动本地稠密检索服务、下载检查点、切换模型与观测验证指标的全部实操细节。一、Search-R1 是什么多轮、工具调用、检索增强Search-R1 是 torchtitan RL 实验torchtitan/experiments/rl内置的一个检索增强生成 RL 示例。它的任务设定非常具体模型拿到一个自然语言问题被赋予一个名为search的工具标准 OpenAI function-calling 工具调用格式模型需要外部事实时发起search工具调用环境把检索到的段落passages以一条tool角色消息回传给模型多轮 rollout 持续进行直到模型停止调用工具、或回合预算耗尽最终答案与黄金答案做精确匹配exact-match, EM得到 0/1 奖励可选地叠加两个把检索行为纳入梯度的调节项。两个设计要点值得先说明思考开关由 renderer 的enable_thinking标志控制而不是在 prompt 里注入think标签。该配方将其设为False任务是短答案事实型问答思维链对 EM 无帮助反而会挤占多轮的 token 预算。如果你的任务受益于推理可在配置中将其翻转为True。整个示例零框架代码它完全运行在框架自带的多轮 rollouterrollout/rollouter.py与连续批处理 generator 之上——示例专属的代码只有本文件夹data / env / rubric / rollouter 四个文件及其配置注册。从源码结构看这套用户写环境、框架驱动循环的分层是刻意设计的MessageEnv基类environment/message.py的 docstring 直接给出了一个计算器环境的实现范式——init()返回初始对话与工具 schemastep()要么回传一条 tool 消息让对话继续要么置doneTrue结束 rollout。Search-R1 的SearchR1Env正是这一范式的落地。二、示例文件夹逐文件解析2.1data.py无限流的 NQ/HotpotQA 数据集SearchR1Dataset 从 HF Hub 数据集PeterJinGo/nq_hotpotqa_train拉取已预处理好的 NQ/HotpotQA parquet无需任何本地预处理对外表现为一个无限、可复现、断点可恢复的样本流行序用seed打乱每次循环wrap重新洗牌保证每个 epoch 都看到新排列。每个样本是冻结 dataclass SearchR1Sample字段含义question模型需要回答的自然语言问题golden_answers可接受的黄金答案字符串列表预测匹配其中任意一个即算 EM 正确数据集的配置项SearchR1Dataset.Configdata.py参数默认值说明filenametrain.parquet从 HF 数据集仓库加载哪个 splittrain.parquet训练或test.parquet验证repo_idPeterJinGo/nq_hotpotqa_train存放预处理后 NQ/HotpotQA parquet 的 HF Hub 数据集仓库data_pathNone本地 parquet 路径设置后覆盖 HF 下载离线场景。parquet 需含question/golden_answers列seed42行序打乱的随机种子data_sourceNone若设置只保留data_source等于该值的行例如nq——合并的 test split 混合了多个数据集None保留全部shuffleTrue用seed打乱行序每次 wrap 重新洗牌。验证时设False使每次验证轮抽取同一批固定样本一个容易被忽视但很有工程价值的细节该数据集实现了state_dict()/load_state_dict()保存 RNG 状态、当前行序与游标位置使一次运行可以在流中途恢复——这是数据集检查点能力的一部分参见 docs/checkpoint.md。2.2env.py定义 search 工具、执行检索、回传 tool 消息SearchR1Env 继承框架的MessageEnv是示例中唯一会做事的逻辑核心有四块1search 工具的 ToolSpec。这是 OpenAI function-calling schema 形式的工具定义renderer 会把它注入 prompt 的 chat template并把模型输出的工具调用解析回completion_message[tool_calls]SEARCH_TOOL: ToolSpec { name: search, description: ( Search a Wikipedia-derived knowledge base and return the top passages for a query. Use it whenever you need external facts to answer the question. ), parameters: { type: object, properties: { query: {type: string, description: The natural-language search query.}, }, required: [query], }, }2指令前缀INSTRUCTION要求模型回答问题时对不确定的事实使用 search 工具并强调最终回复只输出答案本身——实体或短语不带解释或完整句子例如Beijing。这与 EM 奖励严格对应啰嗦的答案会直接失分。3异步检索_search。向检索服务端POST {queries: [query], topk: topk}读回{result: [[{contents: ...}, ...]]}再由_passages_to_string格式化为模型可读的块例如Doc 1(Title: Eiffel Tower) A tower in Paris.错误处理策略值得注意任何传输/服务端异常不抛出而是返回空串并记 warning——一次抖动的请求不会打崩整个 rollout但持续失败也会留下日志而非静默地拿空结果训练。4step()的多轮语义。若本轮completion_message中没有tool_calls说明 assistant 给出了最终答案返回doneTrue结束 rollout若有则用asyncio.gather并发执行所有 search 调用同一轮可能发起多个检索把每个结果包装成{role: tool, content: ...}消息回传doneFalse让对话继续。SearchR1Env.Config的三个参数继承自MessageEnv.Config参数默认值说明search_urlhttp://127.0.0.1:8000/retrieve本地稠密检索服务 URL如换端口/主机在配置中覆盖message_env.search_urltopk3每次查询检索的段落数覆盖message_env.topktimeout_s60.0单次检索请求超时秒回合预算在哪里执行环境本身不数回合——那是外层TokenEnv的职责。environment/token.py 负责把消息空间的环境驱动成 token 空间的 rollout 循环解码 completion、调用 MessageEnv、把环境回复编码回下一轮 prompt并在step()中检查若本轮后num_turns max_num_turns且环境还想继续rollout 以TRUNCATED_MAX_TURNS终止token.py同理max_rollout_tokens会在下一轮 prompt 超出上下文前截断为TRUNCATED_PROMPT_TOO_LONG。此外TokenEnv还统一处理解析失败ERROR_PARSE、长度截断TRUNCATED_LENGTH与 step 超时ERROR_TIMEOUT让 rollout 主循环保持干净。2.3rubric.pyEM 奖励与两个反闭卷作弊调节项RewardExactMatch 的默认行为是纯 EM 0/1最终答案匹配任一黄金答案得 1.0否则 0。匹配前会做归一化——小写、去标点、去掉冠词a/an/the、压缩空白_normalize_answer。判断最终答案有一个边界情形如果 rollout 在最后一次工具调用后结束例如被回合预算截断最后一轮没有文本答案_final_answer返回None即无答案可判得 0 分。配置项RewardExactMatch.Config参数默认值说明score1.0用检索且最终答案正确的得分no_search_penalty0.0从答对但从未调用 search的样本中扣除。0默认 纯 EM设大于 0如 0.2可让闭卷答对的得分低于检索后答对的防止模型靠参数化记忆绕开检索anti closed-book reward hackingretrieval_score0.0最终答案错误/缺失但某次检索曾把黄金答案捞上来作为整词出现在 tool 消息中时的部分得分。0默认 纯 EM打分逻辑rubric.py可以概括为一张表场景得分答对且调用过 searchscore1.0答对但从未 searchscore - no_search_penalty答错/无答案但检索曾捞起黄金答案retrieval_score其余0.0这两个调节项本质上是把检索行为放进奖励梯度前者惩罚不查也会答后者奖励查到了但没答对共同引导策略真正学会使用工具。2.4rollouter.py把数据集 环境 评分组装进框架SearchR1Rollouter 和SearchR1Worker都是纯配置类——全部行为继承自框架的Rollouter/RolloutWorker其设计文档见 rollouter.py 的类 docstringRollouter 之于 rollout 数据正如 Dataloader 之于训练 batch示例只是提供默认配置奖励装配Rubric.Config(reward_fns[RewardExactMatch.Config(weight1.0)], truncation_reward0.0)。truncation_reward0.0意味着被截断、没有最终答案的 rollout 不提供奖励与学习信号token/回合预算TokenEnv.Config(max_rollout_tokens3072, max_num_turns4)——每条 rollout 至多 4 个 assistant 回合、prompt 不超过 3072 token训练数据集SearchR1Dataset.Config(filenametrain.parquet, seed42)验证数据集SearchR1Dataset.Config(filenametest.parquet, seed99, data_sourcenq, shuffleFalse)——只取 NQ split、确定序保证每轮验证抽取同一批留出样本让 EM 曲线可比。三、运行前提数据、检索服务、检查点3.1 数据无需任何准备。NQ/HotpotQA parquet 在首次使用时直接从 HF Hub 数据集PeterJinGo/nq_hotpotqa_train拉取train NQ-test 两个 split。若要改用本地副本把数据集配置的data_path指到一个含question/golden_answers列的 parquet 即可见 2.1 节参数表。3.2 本地稠密检索服务训练之前先启动稠密检索器e5 索引建在 wiki-18 上监听http://127.0.0.1:8000/retrieve。要点是把它固定在空闲 GPU 上避免与 RL 的 GPU 冲突python search-r1/local_dense_retriever/retrieval_server.py \ --index_path $INDEX_PATH/e5_Flat.index \ --corpus_path $CORPUS_PATH/wiki-18.jsonl \ --topk 3 --retriever_name e5 --retriever_model intfloat/e5-base-v2 --faiss_gpu如端口或 topk 不同在配置中覆盖message_env.search_url/message_env.topk对应 env.py 中SearchR1Env.Config的两个字段。3.3 基座检查点下载配置所期望的基座模型。download_hf_assets.py会写入一个以仓库名命名的子目录这正是配置中hf_assets_path所指向的路径python scripts/download_hf_assets.py \ --repo_id meta-models/Muse-Glimmer-30B \ --local_dir torchtitan/experiments/rl/example_checkpoint \ --all把--repo_id换成你配置所选的模型即可例如Qwen/Qwen3-1.7B与 config_registry.py 中各配置写入的hf_assets_path子目录名一致如torchtitan/experiments/rl/example_checkpoint/Qwen3-1.7B。四、启动训练# 示例运行Qwen3-1.7BWB 开启 python torchtitan/experiments/rl/train.py \ --module search_r1 \ --config rl_grpo_qwen3_1_7b_search_r1入口是 train.py基于 Monarch Actor 的分布式训练循环generator 用 vLLM、trainer 用 torchtitan 原生栈分列不同 GPU mesh 并通过 TorchStore 做权重同步。--module search_r1让 ConfigManager 直接发现该示例模块注册的配置入口--config指定 config_registry.py 中的具体配方。训练中盯住validation_reward/_meanNQ test split 上的 EM稳步上升即说明策略正在学会调用search并给出简洁答案。五、配置配方全解析config_registry.pyconfig_registry.py 提供了四个配方全部从示例配置侧完整定义 Search-R1 流程——核心默认值不动其他配置保持 vanilla GRPO 不受影响。以rl_grpo_qwen3_1_7b_search_r1Qwen3-1.7B8 卡4 卡 generator TP4 1 卡 trainer TP1检索服务占剩余 GPU为例关键设置配置块取值说明async_loop500 步每步 8 prompts × 8 samples验证 500 条异步 rollout 循环的规模参数advantageshould_std_normalizeTrueGRPO 组内优势按标准差归一化rendererQwen3RendererConfig(enable_thinkingFalse)关闭思考见第一节优化器 / LRAdamWlr1e-6warmup 2 步linear 衰减min_lr_factor1.0保守的小学习率损失ChunkedLossWrapper(num_chunks8)包裹DAPOLoss(ratio_clip_low0.2, ratio_clip_high0.28)DAPO 式上高下低非对称裁剪无 KL / 无参考模型详见 losses/dapo.py检查点interval50initial_load_in_hfTruelast_save_model_onlyFalsekeep_latest_k3首跑从 HF 加载、重启从 DCP 恢复保留完整非仅模型末次存档保证可恢复keep_latest_k限磁盘generator 采样temperature1.0, top_p1.0, max_tokens512bf16cudagraph 开启vLLM 侧的 rollout 采样参数其余三个配方及差异rl_grpo_qwen3_8b_search_r1与 1.7B 同配方只换模型与 GPU 切分——8 卡 2 卡 generatorTP2 4 卡 trainerTP4fp32 trainer 需要 TP4 才不 OOM。另把 generator 的gpu_memory_limit从默认 0.9 降到 0.6为权重同步的显存尖峰留出空间否则 8B generator 会 OOMrl_grpo_qwen3_30b_a3b_deepep_search_r1_perfQwen3-30B-A3B MoE 的性能配方generator 用 DeepEP v2 cudagraph 路径可跨节点H100 上节点内 NVLink 节点间 IB/RoCEtrainer 保留可反向的 host-synced DeepEP 路径注意 Qwen3-30B-A3B 只有 4 个 KV head所以 generator TP 必须 ≤4trainer 侧使用 FSDP8 × EP8并应用fused_swiglu与helion_rope两个性能 override仅 CUDArl_grpo_muse_glimmer_30b_search_r1Muse Glimmer 30B 配方8 卡 6 卡 trainer FSDP3×TP2 2 卡 generator TP2有两个模型特有的硬约束值得学习Muse Glimmer 只有 2 个 KV headgenerator TP 上限为 2规模扩张要靠 FSDP且必须开 FullAC全激活检查点——因为 Adam 的 m/v 在第一次optimizer.step()才分配第 1→2 步单卡显存会跳升约 8 bytes/param默认的 SelectiveAC 会在第 2 步 OOMFullAC 释放激活内存余量才能扛住。六、训练结果验证集 EM 曲线验证 EM留出 NQ test split贪心解码随策略学会调用search并简洁作答而稳步爬升。由于该配方仍在持续演进下面每条曲线是快照Qwen3-1.7B—— EM 约 0.05 → 约 0.41Qwen3-8B—— EM 约 0.26 → 约 0.45七、无需 GPU、无需检索服务的单元测试想在不启动 vLLM 与检索服务的前提下验证本示例逻辑直接跑 tests/test_search_r1.py 即可——该测试文件对_search做了 monkeypatch用假检索替代网络调用覆盖三类断言环境行为init()暴露且仅暴露search工具、问题进入初始消息带 tool_calls 的step()返回doneFalse且回传一条roletool消息不带 tool_calls 的step()返回doneTrue即最终答案终止健壮性arguments为 JSON 字符串而非 dict时也能正确抽出query——这正是_query_from_tool_call处理的真实解析形态奖励函数用构造的Rollout检索轮 最终答案轮或末轮停在工具调用上模拟截断验证纯 EM、no_search_penalty扣分、retrieval_score部分得分各条路径。八、小结这个示例教会你的扩展模式Search-R1 示例的价值不只在能跑更在于它演示了 torchtitan RL 实验的最小扩展面写一个数据流继承Configurable__iter__吐样本可选state_dict支持恢复写一个MessageEnv子类init给对话与工具 schemastep决定回传工具消息还是结束写一个RewardFn子类对整条多轮 rollout 打分用一个纯配置 Rollouter/Worker 把三者接起来再在config_registry.py里注册一条Controller.Config。多轮循环、token 预算、截断处理、解析失败恢复、连续批处理、权重同步等重活全部由框架层rollouter.py、token.py承担。把它当作模板替换掉检索逻辑与评分规则就能快速派生出其他工具调用型 RL 任务。【免费下载链接】torchtitanA PyTorch native platform for training generative AI models项目地址: https://gitcode.com/GitHub_Trending/to/torchtitan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考