案例

Search-R1

分类:Agentic RL;参考复现总花销 ¥13.30

代码出处与消耗

  • 完整代码来自 KMnO4-zx/llm-agent-rl-lab/03-search-r1,本文核对的代码版本为 52c2f1c
  • 源仓库展示的那次 Evaluation + Training 会话,PyTRIO 合计花销为 ¥13.30。这是一次真实运行的参考值,不是固定报价;实际消耗会随 step 数、轨迹长度、搜索次数和评测规模变化。
  • 训练曲线见 SwanLab 实验记录。由于该案例由多个模块共同组成,本文展示关键代码与完整训练逻辑;可运行的全部文件请以上述源仓库为准。

Search-R1 参考运行的 PyTRIO 消耗

介绍

Search-R1 训练模型在回答问题时自主决定:

  • 什么时候调用搜索;
  • 搜索什么 query;
  • 如何利用搜索结果继续推理;
  • 什么时候停止搜索并输出最终答案。

它和“先检索一次,再把结果塞进 prompt”的普通 RAG 不同。一条 Search-R1 轨迹可能多次交替执行:

assistant 推理
→ search(query)
→ tool observation
→ assistant 继续推理
→ ...
→ Answer: <short answer>

这个复现使用 Qwen/Qwen3.5-4B、PyTRIO 和知乎全局搜索 API,保留 Search-R1 的多轮工具环境、结果奖励、组内相对 advantage、observation token mask 和策略更新,把原论文较重的本地 Wikipedia 检索基础设施替换成在线搜索服务。

需要特别注意:

训练对象是模型的 LoRA 权重。搜索 API 和搜索结果都属于固定环境,不参与训练。

模型学习的是工具调用和答案生成策略,而不是训练搜索引擎。

Search-R1 多轮搜索与组内相对策略更新流程

实验配置

项目本案例配置
Base ModelQwen/Qwen3.5-4B
LoRA rank32
训练框架PyTRIO
搜索环境知乎全局搜索 API,Top 3
训练数据NQ + HotpotQA
每个 step8 道问题
每道问题8 条轨迹
最大搜索次数4
最大 assistant turns6
最大轨迹长度8,192 tokens
RewardExact Match + Format
Advantagereward - group_mean
Lossimportance_sampling

代码结构

Search-R1 不是单文件示例,完整目录由以下模块组成:

03-search-r1/
├── prepare_data.py   # 下载并整理训练集与评测集
├── data.py           # 读取本地 JSONL
├── protocol.py       # 工具 schema、prompt 和 tool-call 解析
├── search.py         # 知乎搜索客户端与调用统计
├── rollout.py        # 多轮工具调用状态机
├── reward.py         # Exact Match + Format reward
├── train.py          # PyTRIO 训练、拆批、更新与 checkpoint
├── eval.py           # Base Model / checkpoint 统一评测
└── analyse.py        # 汇总 checkpoint 评测结果

完整实现请直接查看 源代码目录。下面只展开决定算法行为的主要代码。

环境与数据

本地只需要可联网的 CPU 环境。模型采样、LoRA 前向反向和参数更新由 PyTRIO 远端执行。

git clone https://github.com/KMnO4-zx/llm-agent-rl-lab.git
cd llm-agent-rl-lab

uv sync
trio login
swanlab login

数据来自 Search-R1 公开数据集 PeterJinGo/nq_hotpotqa_train,代码通过固定版本的 ModelScope 镜像下载。执行:

uv run python 03-search-r1/prepare_data.py

脚本会生成:

03-search-r1/datasets/
├── train.jsonl   # 169,615 道 NQ + HotpotQA 训练题
├── dev.jsonl     # 7 个 benchmark 各 10 道,共 70 道
└── test.jsonl    # 完整评测池

训练样本只保留问题和参考答案,不包含人工搜索 query 或标准搜索轨迹:

{
  "id": "...",
  "question": "...",
  "answers": ["..."],
  "data_source": "nq"
}

知乎数据开放平台 申请搜索 API Key 后,复制配置模板:

cp 03-search-r1/.env.example 03-search-r1/.env

然后写入一个或多个 key:

ZHIHU_SEARCH_KEYS=your_first_key,your_second_key,your_third_key

正式训练前应先确认 key 有足够额度,并在训练中持续检查 search/success_ratesearch/error_rate。工具环境不稳定会直接污染 reward。

核心逻辑

1. 把搜索声明为模型工具

工具协议定义在 protocol.py

SEARCH_TOOL = {
    "type": "function",
    "function": {
        "name": "search",
        "description": "Search Zhihu for evidence. Use a concise English query.",
        "parameters": {
            "type": "object",
            "properties": {"query": {"type": "string"}},
            "required": ["query"],
        },
    },
}

prompt_tokens = tokenizer.apply_chat_template(
    messages,
    tools=[SEARCH_TOOL],
    tokenize=True,
    add_generation_prompt=True,
    enable_thinking=False,
)

Qwen3.5 通过 chat template 生成结构化 <tool_call>protocol.py 解析 query,search.py 执行真实搜索,再把标题、内容片段、来源和 URL 作为 role="tool" 的 observation 写回对话。

这里不是依靠 stop word 判断模型是否搜索;停止序列只负责结束当前生成,工具意图来自模型输出的结构化 tool call。

2. 推进多轮、会分叉的轨迹

同一道题的第一轮有共同 prompt,可以一次请求 8 个样本:

sample_async(
    prompt=shared_prompt,
    num_samples=8,
)

第一次搜索后,每条轨迹会产生不同 query 和 observation,下一轮上下文已经分叉:

第一轮:
1 个 shared prompt × num_samples=8

第一次搜索后:
8 个独立 prompt × num_samples=1
多个 sample_async 并发执行

单条轨迹内部必须保持严格因果顺序:

assistant generation
→ search(query)
→ tool observation
→ next assistant generation

不同轨迹之间才可以并发。完整状态机位于 rollout.py

同题轨迹从共享首轮 Prompt 到独立搜索上下文的分叉

3. 只根据最终答案计算 Reward

模型需要在最后输出:

Answer: <short answer>

reward.py 使用三档结果奖励:

最终结果Reward
格式合法且答案正确1.0
格式合法但答案错误0.0
格式非法或没有最终答案-0.1
def score_answer(text: str, references: list[str]) -> RewardResult:
    answer = extract_answer(text)
    if answer is None:
        return RewardResult(-0.1, False, False, None)
    exact_match = any(
        normalize_answer(answer) == normalize_answer(reference)
        for reference in references
    )
    return RewardResult(float(exact_match), True, exact_match, answer)

代码不会奖励搜索次数或中间 query,避免模型为了拿分而无意义地反复搜索。搜索路径由策略自己探索,最终答案负责提供 outcome reward。

4. 在完整同题 group 内计算 Advantage

每道题的 8 条轨迹全部结束后,先计算:

Ai=rimean(r1,r2,,r8)A_i = r_i - \operatorname{mean}(r_1, r_2, \ldots, r_8)

对应代码:

mean_reward = sum(item.reward for item in group) / len(group)
for item in group:
    item.advantage = item.reward - mean_reward

如果同组 reward 完全相同,所有 advantage 都是 0,代码会跳过该组。不能先把轨迹随意拆成 micro-batch,再在小 batch 内重新计算均值,否则已经改变 group-relative 算法。

5. 搜索结果进入上下文,但不进入 Loss

一条轨迹中同时存在模型动作和环境 observation:

system / user / tool observation  → advantage = 0
assistant tool call / final answer → advantage = trajectory_advantage

Search-R1 中 reward、advantage 与 token loss mask 的关系

train.pybuild_datum() 会把多轮前缀拼成一条连续轨迹。工具返回的 token 保留为后续推理上下文,但用零 old logprob 和零 advantage 屏蔽:

full_tokens.extend(delta_observation)
full_tokens.extend(turn.completion_tokens)

old_logprobs_by_token.extend([0.0] * len(delta_observation))
old_logprobs_by_token.extend(turn.logprobs)

advantages_by_token.extend([0.0] * len(delta_observation))
advantages_by_token.extend(
    [trajectory.advantage] * len(turn.completion_tokens)
)

input_tokens = full_tokens[:-1]
target_tokens = full_tokens[1:]
old_logprobs = old_logprobs_by_token[1:]
advantages = advantages_by_token[1:]

datum = trio.Datum(
    model_input=trio.ModelInput.from_ints(input_tokens),
    loss_fn_inputs={
        "target_tokens": np.asarray(target_tokens, dtype=np.int64),
        "logprobs": np.asarray(old_logprobs, dtype=np.float32),
        "advantages": np.asarray(advantages, dtype=np.float32),
    },
)

四个字段必须经过同一次自回归右移并严格等长。非零 old logprob 必须来自 rollout 当时的 Student sampler,不能在参数更新后重新计算。

6. 完整算完 group,再拆 micro-batch 更新

一个逻辑 step 最多产生:

8 questions × 8 trajectories = 64 trajectories

每条轨迹又允许达到 8,192 tokens,因此训练代码先完成整个 rollout batch 的 reward 和 advantage,再按 padding 后的矩形大小动态装箱:

完整 rollout group
→ reward
→ group-relative advantage
→ 每条完整轨迹构造 Datum
→ 拆成 micro-batch
→ 多次累积 forward/backward
→ 整个逻辑 step 只做一次 optimizer step

当前限制是单条 Datum 不超过 8,192 tokens、单个 micro-batch 不超过 32 个 Datum,并且 items × max_sequence_length 不超过 64,000。

trajectories = rollout_batch(...)
datums = build_training_datums(trajectories)
micro_batches = pack_micro_batches(datums)

for micro_batch in micro_batches:
    training_client.forward_backward(
        weight_micro_batch_for_global_mean(
            micro_batch,
            total_samples=len(trajectories),
        ),
        loss_fn="importance_sampling",
    ).result()

if micro_batches:
    training_client.optim_step(adam_params).result()

由于远端对每次 forward_backward 的样本取均值,不同大小的 micro-batch 会按 n_k / N 缩放 advantage,使梯度累积仍等价于完整 logical batch 的全局样本均值。

运行训练

建议先跑 20 step 验证完整链路:

uv run python 03-search-r1/train.py \
    --max-steps 20 \
    --questions-per-batch 8 \
    --group-size 8 \
    --save-every 5 \
    --swanlab-mode online

每 5 step 会保存两份权重:

*-state    # 包含优化器状态,用于断点续训
*-weights  # 用于采样和评测

20 step 主要用于确认:

  • reward/format 是否开始上升;
  • 模型是否会在多轮搜索后稳定输出 Answer:
  • reward/correct 是否出现变化;
  • search/success_rate 是否稳定;
  • rollout/degenerate_group_rate 是否过高。

评测

先用相同环境评测 Base Model:

uv run python 03-search-r1/eval.py \
    --batch-size 16 \
    --output 03-search-r1/eval_result/eval_results.jsonl

再把训练日志中保存的 sampler weights 路径传给 evaluator:

uv run python 03-search-r1/eval.py \
    --batch-size 16 \
    --model-path 'trio://YOUR_STEP_20_SAMPLER_WEIGHTS' \
    --output 03-search-r1/eval_result/eval_results_rl_step_20.jsonl

搜索额度有限时,可以先增加 --limit 20 做链路检查,再运行固定 70 题评测。

参考运行在固定 70 题上的结果为:

Search-R1 Base Model 与各 checkpoint 的 Macro EM 和 Format Rate

Base Model:  Macro EM 28.57% · Format 58.57%
RL Step 20:  Macro EM 31.43% · Format 87.14%
RL Step 50:  Macro EM 45.71% · Format 94.29%

这里的 Step 20 来自较早的小规模 run,实时搜索条件与主实验不完全相同;70 题评测也更适合验证行为变化,而不是直接作为论文级结论。

训练时重点观察

指标作用
reward/mean整条轨迹的平均结果
reward/correct最终答案正确率
reward/format是否学会结束搜索并按格式回答
rollout/search_calls每条轨迹平均搜索次数
rollout/turns多轮轨迹长度变化
rollout/degenerate_group_rate没有组内相对信号的问题比例
train/loss_tokens_per_rollout_batch实际参与 loss 的 assistant token 数
search/success_rate工具环境是否可靠
search/error_ratereward 是否可能被搜索错误污染
search/latency搜索是否成为 rollout 瓶颈

如果 reward/correct 下跌的同时 search/error_rate 上升,应先检查搜索环境,不能直接判断模型能力退化。

复现边界

这个案例复现的是 Search-R1 的核心训练闭环,不是原论文基础设施和分数的逐项复刻:

  • 模型、搜索后端、训练框架和部分配置与原论文不同;
  • 当前固定评测集只有 70 道题;
  • 在线搜索会受到额度、超时和结果变化影响;
  • 搜索后端保持固定,训练的是模型的工具使用与回答策略。

完整实现、实验图和更详细的结果讨论请阅读源仓库的 Search-R1 README

这篇文档对你有帮助吗?

本页目录