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。 - 整个示例"零框架代码":它完全运行在框架自带的多轮 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_path | None | 本地 parquet 路径;设置后覆盖 HF 下载(离线场景)。parquet 需含question/golden_answers列 |
seed | 42 | 行序打乱的随机种子 |
data_source | None | 若设置,只保留data_source等于该值的行(例如"nq")——合并的 test split 混合了多个数据集;None保留全部 |
shuffle | True | 用seed打乱行序,每次 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 |
topk | 3 | 每次查询检索的段落数;覆盖message_env.topk |
timeout_s | 60.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):
| 参数 | 默认值 | 说明 |
|---|---|---|
score | 1.0 | 用检索且最终答案正确的得分 |
no_search_penalty | 0.0 | 从"答对但从未调用 search"的样本中扣除。0(默认)= 纯 EM;设大于 0(如 0.2)可让闭卷答对的得分低于检索后答对的,防止模型靠参数化记忆绕开检索(anti closed-book reward hacking) |
retrieval_score | 0.0 | 最终答案错误/缺失,但某次检索曾把黄金答案"捞上来"(作为整词出现在 tool 消息中)时的部分得分。0(默认)= 纯 EM |
打分逻辑(rubric.py)可以概括为一张表:
| 场景 | 得分 |
|---|---|
| 答对且调用过 search | score(1.0) |
| 答对但从未 search | score - 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_loop | 500 步;每步 8 prompts × 8 samples;验证 500 条 | 异步 rollout 循环的规模参数 |
advantage | should_std_normalize=True | GRPO 组内优势按标准差归一化 |
renderer | Qwen3RendererConfig(enable_thinking=False) | 关闭思考(见第一节) |
| 优化器 / LR | AdamWlr=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=50,initial_load_in_hf=True,last_save_model_only=False,keep_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_swiglu与helion_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 实验的最小扩展面:
- 写一个数据流(继承
Configurable,__iter__吐样本,可选state_dict支持恢复); - 写一个
MessageEnv子类(init给对话与工具 schema,step决定回传工具消息还是结束); - 写一个
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),仅供参考