news 2026/9/17 5:02:50

torchtitan Search-R1 示例:用 search 工具构建多轮检索增强 GRPO 训练流水线的完整指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
torchtitan Search-R1 示例:用 search 工具构建多轮检索增强 GRPO 训练流水线的完整指南

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 奖励,可选地叠加两个"把检索行为纳入梯度"的调节项。

两个设计要点值得先说明:

  1. 思考开关由 renderer 的enable_thinking标志控制,而不是在 prompt 里注入think标签。该配方将其设为False:任务是短答案事实型问答,思维链对 EM 无帮助,反而会挤占多轮的 token 预算。如果你的任务受益于推理,可在配置中将其翻转为True
  2. 整个示例"零框架代码":它完全运行在框架自带的多轮 rollouter(rollout/rollouter.py)与连续批处理 generator 之上——示例专属的代码只有本文件夹(data / env / rubric / rollouter 四个文件)及其配置注册。

从源码结构看,这套"用户写环境、框架驱动循环"的分层是刻意设计的:MessageEnv基类(environment/message.py)的 docstring 直接给出了一个计算器环境的实现范式——init()返回初始对话与工具 schema,step()要么回传一条 tool 消息让对话继续,要么置done=True结束 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.Config,data.py):

参数默认值说明
filename"train.parquet"从 HF 数据集仓库加载哪个 split:train.parquet(训练)或test.parquet(验证)
repo_id"PeterJinGo/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保留全部
shuffleTrueseed打乱行序,每次 wrap 重新洗牌。验证时设False,使每次验证轮抽取同一批固定样本

一个容易被忽视但很有工程价值的细节:该数据集实现了state_dict()/load_state_dict()(保存 RNG 状态、当前行序与游标位置),使一次运行可以在流中途恢复——这是数据集检查点能力的一部分(参见 docs/checkpoint.md)。

2.2env.py:定义 search 工具、执行检索、回传 tool 消息

SearchR1Env 继承框架的MessageEnv,是示例中唯一"会做事"的逻辑,核心有四块:

(1)search 工具的 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,但持续失败也会留下日志而非静默地拿空结果训练。

(4)step()的多轮语义。若本轮completion_message中没有tool_calls,说明 assistant 给出了最终答案,返回done=True结束 rollout;若有,则用asyncio.gather并发执行所有 search 调用(同一轮可能发起多个检索),把每个结果包装成{"role": "tool", "content": ...}消息回传,done=False让对话继续。

SearchR1Env.Config的三个参数(继承自MessageEnv.Config):

参数默认值说明
search_url"http://127.0.0.1:8000/retrieve"本地稠密检索服务 URL;如换端口/主机在配置中覆盖message_env.search_url
topk3每次查询检索的段落数;覆盖message_env.topk
timeout_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.py:EM 奖励与两个"反闭卷作弊"调节项

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 hacking)
retrieval_score0.0最终答案错误/缺失,但某次检索曾把黄金答案"捞上来"(作为整词出现在 tool 消息中)时的部分得分。0(默认)= 纯 EM

打分逻辑(rubric.py)可以概括为一张表:

场景得分
答对且调用过 searchscore(1.0)
答对但从未 searchscore - no_search_penalty
答错/无答案,但检索曾捞起黄金答案retrieval_score
其余0.0

这两个调节项本质上是把"检索行为"放进奖励梯度:前者惩罚"不查也会答",后者奖励"查到了但没答对",共同引导策略真正学会使用工具。

2.4rollouter.py:把数据集 + 环境 + 评分组装进框架

SearchR1Rollouter 和SearchR1Worker都是纯配置类——全部行为继承自框架的Rollouter/RolloutWorker(其设计文档见 rollouter.py 的类 docstring:Rollouter 之于 rollout 数据,正如 Dataloader 之于训练 batch),示例只是提供默认配置:

  • 奖励装配Rubric.Config(reward_fns=[RewardExactMatch.Config(weight=1.0)], truncation_reward=0.0)truncation_reward=0.0意味着被截断、没有最终答案的 rollout 不提供奖励与学习信号;
  • token/回合预算TokenEnv.Config(max_rollout_tokens=3072, max_num_turns=4)——每条 rollout 至多 4 个 assistant 回合、prompt 不超过 3072 token;
  • 训练数据集SearchR1Dataset.Config(filename="train.parquet", seed=42)
  • 验证数据集SearchR1Dataset.Config(filename="test.parquet", seed=99, data_source="nq", shuffle=False)——只取 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.7B),W&B 开启 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/_mean(NQ test split 上的 EM)稳步上升,即说明策略正在学会调用search并给出简洁答案。

五、配置配方全解析(config_registry.py)

config_registry.py 提供了四个配方,全部"从示例配置侧"完整定义 Search-R1 流程——核心默认值不动,其他配置保持 vanilla GRPO 不受影响。以rl_grpo_qwen3_1_7b_search_r1(Qwen3-1.7B,8 卡:4 卡 generator TP=4 + 1 卡 trainer TP=1,检索服务占剩余 GPU)为例,关键设置:

配置块取值说明
async_loop500 步;每步 8 prompts × 8 samples;验证 500 条异步 rollout 循环的规模参数
advantageshould_std_normalize=TrueGRPO 组内优势按标准差归一化
rendererQwen3RendererConfig(enable_thinking=False)关闭思考(见第一节)
优化器 / LRAdamWlr=1e-6;warmup 2 步,linear 衰减,min_lr_factor=1.0保守的小学习率
损失ChunkedLossWrapper(num_chunks=8)包裹DAPOLoss(ratio_clip_low=0.2, ratio_clip_high=0.28)DAPO 式"上高下低"非对称裁剪;无 KL / 无参考模型,详见 losses/dapo.py
检查点interval=50initial_load_in_hf=Truelast_save_model_only=Falsekeep_latest_k=3首跑从 HF 加载、重启从 DCP 恢复;保留完整(非仅模型)末次存档保证可恢复,keep_latest_k限磁盘
generator 采样temperature=1.0, top_p=1.0, max_tokens=512,bf16,cudagraph 开启vLLM 侧的 rollout 采样参数

其余三个配方及差异:

  • rl_grpo_qwen3_8b_search_r1:与 1.7B 同配方,只换模型与 GPU 切分——8 卡 = 2 卡 generator(TP=2)+ 4 卡 trainer(TP=4,fp32 trainer 需要 TP=4 才不 OOM)。另把 generator 的gpu_memory_limit从默认 0.9 降到 0.6,为权重同步的显存尖峰留出空间(否则 8B generator 会 OOM);
  • rl_grpo_qwen3_30b_a3b_deepep_search_r1_perf:Qwen3-30B-A3B MoE 的性能配方,generator 用 DeepEP v2 cudagraph 路径(可跨节点,H100 上节点内 NVLink + 节点间 IB/RoCE),trainer 保留可反向的 host-synced DeepEP 路径;注意 Qwen3-30B-A3B 只有 4 个 KV head,所以 generator TP 必须 ≤4;trainer 侧使用 FSDP=8 × EP=8,并应用fused_swigluhelion_rope两个性能 override(仅 CUDA);
  • rl_grpo_muse_glimmer_30b_search_r1:Muse Glimmer 30B 配方(8 卡 = 6 卡 trainer FSDP3×TP2 + 2 卡 generator TP2),有两个模型特有的硬约束值得学习:Muse Glimmer 只有 2 个 KV head,generator TP 上限为 2,规模扩张要靠 FSDP;且必须开 FullAC(全激活检查点)——因为 Adam 的 m/v 在第一次optimizer.step()才分配,第 1→2 步单卡显存会跳升约 8 bytes/param,默认的 SelectiveAC 会在第 2 步 OOM,FullAC 释放激活内存余量才能扛住。

六、训练结果:验证集 EM 曲线

验证 EM(留出 NQ test split,贪心解码)随策略学会调用search并简洁作答而稳步爬升。由于该配方仍在持续演进,下面每条曲线是"快照":

Qwen3-1.7B—— EM 约 0.05 → 约 0.41

Qwen3-8B—— EM 约 0.26 → 约 0.45

七、无需 GPU、无需检索服务的单元测试

想在不启动 vLLM 与检索服务的前提下验证本示例逻辑,直接跑 tests/test_search_r1.py 即可——该测试文件对_search做了 monkeypatch(用假检索替代网络调用),覆盖三类断言:

  • 环境行为init()暴露且仅暴露search工具、问题进入初始消息;带 tool_calls 的step()返回done=False且回传一条role="tool"消息;不带 tool_calls 的step()返回done=True(即最终答案终止);
  • 健壮性arguments为 JSON 字符串(而非 dict)时也能正确抽出query——这正是_query_from_tool_call处理的真实解析形态;
  • 奖励函数:用构造的Rollout(检索轮 + 最终答案轮,或末轮停在工具调用上模拟截断)验证纯 EM、no_search_penalty扣分、retrieval_score部分得分各条路径。

八、小结:这个示例教会你的扩展模式

Search-R1 示例的价值不只在"能跑",更在于它演示了 torchtitan RL 实验的最小扩展面

  1. 写一个数据流(继承Configurable__iter__吐样本,可选state_dict支持恢复);
  2. 写一个MessageEnv子类(init给对话与工具 schema,step决定回传工具消息还是结束);
  3. 写一个RewardFn子类,对整条多轮 rollout 打分;
  4. 用一个纯配置 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),仅供参考

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

WiFi图标消失不用慌:从软件到硬件的完整修复指南

说实话&#xff0c;干了这么多年装机维护&#xff0c;遇到最多的情况之一就是“网络重置后WiFi图标不见了”或者“电脑恢复出厂后无线网络直接消失”。这问题看着小&#xff0c;真碰上的时候非常折腾人&#xff0c;尤其是在急着联网干活的时候&#xff0c;网线一拔、图标一消失…

作者头像 李华
网站建设 2026/9/17 5:01:39

Dynamics 365 FO 建表全指南:从AOT到数据库同步的完整流程

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/17 5:01:35

DPABI fMRI预处理中NIfTI头文件写入错误的解决方案

1. 问题现象与背景解析最近在使用DPABI进行fMRI数据预处理时&#xff0c;不少同行遇到了一个典型报错&#xff1a;"错误使用 nifti/create (line 26) Unable to write header for..."。这个错误通常发生在协变量分析阶段&#xff0c;表现为程序突然中断并弹出红色错误…

作者头像 李华
网站建设 2026/9/17 5:00:13

AD域管理升级实战:从脚本运维到可审计可追溯的企业级运营

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/17 4:59:54

Buck电路滑模控制设计与Simulink仿真实践

1. 项目背景与核心价值Buck电路作为电力电子领域最基础的DC-DC降压拓扑&#xff0c;在电源适配器、车载供电、工业控制等领域应用广泛。但传统PID控制在负载突变或输入电压波动时容易出现超调、振荡等问题。去年我在设计一款医疗设备电源模块时&#xff0c;就遇到过输出纹波超标…

作者头像 李华
网站建设 2026/9/17 4:58:18

游戏服务器选型全指南:业务拆解、硬件指标与配置方案

我接触过不少从零起步的游戏项目&#xff0c;发现一个挺有意思的现象&#xff1a;很多团队在讨论玩法、美术、程序架构时非常投入&#xff0c;唯独到了服务器选型这一步&#xff0c;草率得很。要么直接复制网上所谓“标配”&#xff0c;要么干脆挑个最便宜的云主机先跑起来再说…

作者头像 李华