目录/第八章 Agentic-RL

第八章 大模型强化学习

在第四章中,我们从人类反馈强化学习(Reinforcement Learning from Human Feedback,RLHF)出发,介绍了奖励模型、近端策略优化(Proximal Policy Optimization,PPO)以及指令对齐的基本流程。第六章进一步给出了大模型训练的工程基础,第七章则让模型拥有了规划、记忆与调用工具的能力。

当 Agent 真正进入搜索、代码执行等环境后,仅靠监督微调已经很难覆盖所有可能的交互轨迹。模型需要自己尝试动作,根据最终结果判断一条轨迹是否有效,再把成功经验更新到策略中。大模型强化学习由此从“让回答更符合偏好”,逐步走向“让模型在环境中学会行动”。

本章围绕四个彼此衔接的主题展开。首先介绍当前大模型强化学习中使用广泛的 Group Relative Policy Optimization(GRPO),建立 rollout、reward、advantage 与策略更新的完整认识;随后介绍 On-Policy Distillation(OPD),观察 Student 如何在自己的状态分布上持续获得 Teacher 的逐 token 指导;最后以 Search-R1 和 ReTool 为例,把同一套训练框架扩展到搜索引擎与代码解释器,完成两条可以运行的 Agentic RL 实践。

第八章内容结构

图 8.1 第八章内容结构

本章全部示例均使用 PyTRIO 0.2.6。PyTRIO 将本地的数据处理、环境交互和训练控制,与远端的模型采样、前向反向传播和权重管理连接起来。四个示例虽然任务不同,核心链路始终可以概括为:

  1. 使用当前策略生成一条或一组轨迹;
  2. 根据结果或 Teacher 反馈计算训练信号;
  3. 将 prompt、response、旧策略 logprob 和 advantage 对齐为 Datum
  4. 调用 forward_backward()optim_step() 更新策略;
  5. 刷新采样客户端,让下一轮 rollout 使用新权重。

配套代码位于 docs/chapter8。从仓库根目录安装公共依赖:

uv venv --python 3.13
source .venv/bin/activate
uv pip install -r docs/chapter8/requirements.txt
trio login

PyTRIO 服务支持的模型会持续更新,正式实验前可以使用下面的代码查看当前列表:

import pytrio as trio

service_client = trio.ServiceClient()
print(service_client.get_supported_models())

扩展阅读: 如果你希望继续了解更多 Agentic RL 算法与可运行实践,可以阅读作者维护的另一个开源项目 agentic-rl-lab。除本章介绍的内容外,该仓库还包含 OPSD、DAPO、GSPO、ALFWorld 等算法与环境的原理拆解、训练代码和实验记录。

注:作者选择 PyTrio 而不是 Verl 作为本章示例,意在希望降低 Agentic RL 的工程门槛与实验成本。按照文中的最小试跑配置,作者力争将四个示例的基础实验预算控制在 200 元以内,让读者亲手跑通 rollout、reward、advantage 与策略更新闭环。实际费用会随模型、采样长度、训练步数和评测规模变化。

8.1 GRPO

代码来源: 本节依据作者开源仓库 agentic-rl-lab/01-grpo 中的 01-demo-sync.py02-demo-async.py 整理。正文逐段讲解同步版 01-demo-sync.py,异步版作为完整配套实现保留。

GRPO 最早在 DeepSeekMath 中被系统提出。它保留了 PPO 的 on-policy 更新思想,同时使用同一道题的一组采样结果估计相对优势,从而省去单独训练 Value Model 的过程。对于具有可验证答案的数学、代码和逻辑任务,这种方法能够直接利用结果奖励,已经成为大模型强化学习最常见的基础算法之一。

8.1.1 从语言模型到强化学习策略

自回归语言模型接收 prompt $x$,并依次生成 token 序列 $y=(y_1,y_2,\ldots,y_T)$。从概率建模的角度看,完整回答的概率可以写为:

$$\pi_\theta(y\mid x) = \prod_{t=1}^{T} \pi_\theta(y_t\mid x,y_{\lt t}),$$

其中, $\pi_\theta$ 表示参数为 $\theta$ 的语言模型策略。将语言模型放入强化学习框架后,可以得到表 8.1 所示的对应关系。

强化学习概念大语言模型中的含义
状态 $s_t$prompt 与已经生成的 token
动作 $a_t$下一个 token
策略 $\pi_\theta(a_t\mid s_t)$模型对下一个 token 的概率分布
轨迹 $\tau$一段完整回答,或包含工具调用的多轮交互
环境数据集判题器、搜索引擎、代码解释器等
奖励 $R(\tau)$答案正确性、格式、工具执行结果等反馈

一次模型生成通常称为一次 rollout。训练的目标,是提高高奖励轨迹中动作出现的概率,同时降低低奖励轨迹中动作出现的概率。最直观的策略梯度形式为:

$$\nabla_\theta J(\theta) = \mathbb{E}{\tau\sim\pi\theta} \left[ A(\tau) \nabla_\theta \log \pi_\theta(\tau) \right],$$

其中, $A(\tau)$ 是优势函数。优势为正,说明这条轨迹相对值得鼓励;优势为负,说明当前策略需要减少类似行为。

大模型的动作空间是整个词表,一段回答又可能包含数百到数千个 token。训练时还要控制新旧策略的变化幅度,防止一次更新破坏模型已有能力。因此,真正困难的部分通常集中在两个问题上:如何得到稳定的 advantage,以及如何安全地使用这些 advantage 更新策略。

8.1.2 从 PPO 到 GRPO

PPO 通常使用一个 Value Model 估计状态价值 $V_\phi(s_t)$,再结合回报计算 advantage。大模型场景中的 Value Model 往往与策略模型规模接近,它会带来额外的显存、训练和同步成本。价值估计误差也会进一步传递到策略更新。

GRPO 改用组内相对比较。对于一个问题 $x$,先从旧策略 $\pi_{\theta_{\mathrm{old}}}$ 中采样 $G$ 个回答:

$${y_1,y_2,\ldots,y_G} \sim \pi_{\theta_{\mathrm{old}}}(\cdot\mid x).$$

判题器分别给出奖励 $r_1,r_2,\ldots,r_G$。经典 GRPO 使用组内均值和标准差对奖励进行标准化:

$$A_i = \frac{r_i-\mathrm{mean}(r_1,\ldots,r_G)}{\mathrm{std}(r_1,\ldots,r_G)+\varepsilon}.$$

这样一来,同一道题中的高分回答获得正 advantage,低分回答获得负 advantage。问题本身的难度被组内基线抵消,策略可以更专注地学习“在相同条件下,哪些生成方式更好”。

GRPO 的分组采样与策略更新

图 8.2 GRPO 的分组采样与策略更新

例如,对同一道数学题采样 4 个回答,规则判题器得到奖励 $[1,0,0,1]$。组内均值为 $0.5$,只做均值中心化时,4 个 advantage 分别为:

$$[0.5,-0.5,-0.5,0.5].$$

两个正确回答会被鼓励,两个错误回答会被抑制。本章代码为了直接展示组内基线,默认采用:

$$A_i=r_i-\bar r.$$

完整实验也可以加入标准差归一化。两种写法表达了相同的组内相对思想,但梯度尺度不同,比较实验时应固定 advantage 的计算方式。

这里还有一个容易忽略的退化情形。当一组回答全部正确,或者全部错误时,每个 $A_i$ 都等于 0,这组数据不会产生有效策略梯度。若训练日志中退化组比例长期很高,通常需要调整题目难度、采样温度、基础模型能力或 group size。

8.1.3 奖励与可验证强化学习

GRPO 只负责把组内奖励转换为相对优势,奖励质量仍然决定了模型最终学到什么。数学和代码任务通常能够使用确定性的规则判题器,这类训练也被称为 Reinforcement Learning with Verifiable Rewards(RLVR)。

本章以 GSM8K 为例,要求模型将最终数值写在 \boxed{} 中。最小判题逻辑如下:

def extract_boxed(text: str) -> str | None:
    matches = re.findall(r"\\boxed\{([^}]+)\}", text)
    if not matches:
        return None
    return matches[-1].strip()


def normalize_answer(text: str) -> str:
    return text.replace(",", "").strip().rstrip(".")


def grade_answer(response: str, ground_truth: str) -> float:
    answer = extract_boxed(response)
    if answer is None:
        return 0.0
    return 1.0 if normalize_answer(answer) == normalize_answer(ground_truth) else 0.0

规则奖励简单、便宜,并且可以重复计算。不过,规则一旦存在漏洞,模型就可能学习利用漏洞获得高分。设计 reward 时应至少检查以下内容:

  • 答案抽取是否稳健:同一个正确答案可能出现整数、小数、分数或带逗号数字等形式;
  • 格式奖励是否压过正确性:格式只能作为辅助约束,不能让格式正确的错误答案获得主要奖励;
  • 参考答案是否可靠:错误标签会给模型提供方向相反的训练信号;
  • 判题器是否泄露信息:模型不应通过 prompt 或工具输出直接读到标准答案;
  • 奖励是否具有区分度:大量全 0 或全 1 的组都会让 GRPO 失去相对比较信号。

正式训练前,可以先把基础模型 rollout 保存下来,离线人工检查一批“回答—抽取结果—reward”三元组。这个步骤往往比直接调学习率更早发现问题。

8.1.4 策略比率与更新损失

rollout 由旧策略生成,参数更新发生在当前策略上。对第 $i$ 条回答的第 $t$ 个 token,定义新旧策略概率比:

$$\rho_{i,t}(\theta) = \frac{\pi_\theta(y_{i,t}\mid x,y_{i,\lt t})}{\pi_{\theta_{\mathrm{old}}}(y_{i,t}\mid x,y_{i,\lt t})} = \exp\left(\log\pi_\theta-\log\pi_{\theta_{\mathrm{old}}}\right).$$

最直接的 importance sampling 目标可以写为:

$$J_{\mathrm{IS}}(\theta) = \mathbb{E}{i,t} \left[\rho{i,t}(\theta)A_i\right].$$

如果概率比偏离 1 太远,少数 token 可能主导梯度。PPO 使用截断目标限制单次更新幅度:

$$J_{\mathrm{PPO}}(\theta) = \mathbb{E}{i,t} \left[\min\left(\rho{i,t}A_i, \mathrm{clip}(\rho_{i,t},1-\epsilon,1+\epsilon)A_i\right)\right].$$

DeepSeekMath 中的经典 GRPO 目标采用了 PPO 风格的 clipping,并加入与参考策略之间的 KL 约束。本章的同步版 01-demo-sync.py 默认选择 PyTRIO 内置的 importance_sampling,便于观察 old_logprobsadvantages 和当前策略之间的数据关系;将命令行参数改为 --loss-fn ppo,即可使用内置 PPO loss。

因此,GRPO 与 PPO 需要放在不同层次理解:

  • GRPO 描述如何对同题分组采样,并由组内奖励得到 advantage;
  • importance sampling 或 PPO 描述如何把 advantage 转换为策略梯度。

8.1.5 从头实现同步版 GRPO

本节按照同步版脚本的执行顺序,把一份完整的 GRPO 训练代码拆开说明。配套目录提供两个可以独立运行的文件:

文件执行方式用途
01-demo-sync.py同步正文逐段讲解的教学主线
02-demo-async.py异步并发处理同一 batch 中不同 prompt 的 rollout

两个文件使用相同的数据、reward、advantage、Datum 和 loss。下面只展开同步版,读者先沿着线性的控制流理解一次完整更新。

(1)定义训练配置与 rollout 数据

脚本只保留 PyTRIO 直接提供的 importance_samplingppo 两种策略更新 loss:

LOSS_FNS = ("importance_sampling", "ppo")

GRPOConfig 集中保存模型、采样、优化器和 SwanLab 参数。RolloutSample 保存一条 completion 后续训练需要的全部信息:

@dataclass
class RolloutSample:
    """一条采样结果,以及构造 importance_sampling 所需的旧策略 logprobs。"""

    tokens: list[int]
    logprobs: list[float]
    text: str
    reward: float
    advantage: float

这里的 logprobs 必须来自 rollout 时的 Student。策略更新之后重新计算得到的是新策略 logprob,不能替代旧策略概率。parse_args() 将命令行参数转换为 GRPOConfig,并通过 choices=LOSS_FNS 限定只使用这两个内置 loss。

(2)加载 GSM8K 并构造 prompt

GSM8K 的标准答案位于 #### 后面。模型输出则要求把最终数字写在 \boxed{} 中。代码对两侧答案做相同的轻量归一化:

def extract_boxed(text: str) -> str | None:
    """取最后一个 \\boxed{...} 作为模型最终答案。"""
    matches = re.findall(r"\\boxed\{([^}]+)\}", text)
    if not matches:
        return None
    return matches[-1].strip()


def normalize_answer(text: str) -> str:
    """GSM8K 答案只做轻量归一化,避免 1,000 和 1000 被判成不同。"""
    return text.replace(",", "").strip().rstrip(".")


def grade_answer(response: str, ground_truth: str) -> float:
    """boxed answer 与标准答案完全一致时给 1,否则给 0。"""
    answer = extract_boxed(response)
    if answer is None:
        return 0.0
    return 1.0 if normalize_answer(answer) == normalize_answer(ground_truth) else 0.0


def extract_gsm8k_answer(answer_text: str) -> str:
    """GSM8K 的最终答案位于 `####` 后面。"""
    match = re.search(r"####\s*(.+)", answer_text)
    if match is None:
        raise ValueError(f"No GSM8K final answer found: {answer_text!r}")
    return normalize_answer(match.group(1))

build_prompt() 将一个 few-shot 示例与当前问题交给模型自己的 chat template,再编码成 token:

def build_prompt(tokenizer: Any, question: str) -> list[int]:
    messages = [
        *FEWSHOT_PREFIX,
        {"role": "user", "content": question + QUESTION_SUFFIX},
    ]
    prompt_text = tokenizer.apply_chat_template(
        messages,
        tokenize=False,
        add_generation_prompt=True,
        enable_thinking=False,
    )
    prompt_tokens = tokenizer.encode(prompt_text, add_special_tokens=False)
    if not prompt_tokens:
        raise ValueError("Prompt tokens are empty")
    return prompt_tokens

训练集通过 load_dataset("openai/gsm8k", "main", split="train") 加载。pick_batch() 在小步试跑时允许数据回绕;打开 --all-data 后,每条训练样本最多使用一次。

(3)对同一道题采样一组回答

run_rollout_group() 是 GRPO rollout 的核心。一次 sample() 接收一个 prompt,并通过 num_samples=group_size 返回同题的一组 completion:

def run_rollout_group(
    sampling_client: Any,
    tokenizer: Any,
    prompt_tokens: list[int],
    ground_truth: str,
    sampling_params: trio.SamplingParams,
    group_size: int,
) -> list[RolloutSample]:
    result = sampling_client.sample(
        prompt=trio.ModelInput.from_ints(prompt_tokens),
        num_samples=group_size,
        sampling_params=sampling_params,
        return_text=True,
    ).result()

    rewards: list[float] = []
    raw_samples: list[tuple[list[int], list[float], str]] = []

    for sequence in result.sequences:
        text = sequence.text
        if text is None:
            text = tokenizer.decode(sequence.tokens, skip_special_tokens=True)

        tokens = list(sequence.tokens)
        logprobs = [float(value) for value in sequence.logprobs]
        if len(tokens) != len(logprobs):
            raise ValueError(
                f"Generated token/logprob length mismatch: "
                f"{len(tokens)} != {len(logprobs)}"
            )

        reward = grade_answer(text, ground_truth)
        rewards.append(reward)
        raw_samples.append((tokens, logprobs, text))

    mean_reward = sum(rewards) / len(rewards)
    return [
        RolloutSample(
            tokens=tokens,
            logprobs=logprobs,
            text=text,
            reward=reward,
            advantage=reward - mean_reward,
        )
        for (tokens, logprobs, text), reward
        in zip(raw_samples, rewards, strict=True)
    ]

这段代码依次完成四件事:

  1. 保存 completion token;
  2. 保存采样时旧策略对每个 completion token 的 logprob;
  3. 使用规则判题器得到 0 或 1 的 reward;
  4. 计算 advantage = reward - mean_reward

同组回答全部正确或全部错误时,每条 advantage 都为 0。主循环会把这种 group 记为退化组并跳过,避免提交没有训练信号的 Datum。

(4)构造右移对齐的 PyTRIO Datum

自回归模型用当前位置 token 预测下一个 token。设 prompt 长度为 $m$,completion 长度为 $T$,脚本构造的四个数组长度均为 $m+T-1$:

def build_grpo_datum(
    prompt_tokens: list[int],
    sample: RolloutSample,
) -> trio.Datum:
    if not sample.tokens:
        raise ValueError("Cannot train on an empty completion")

    observation_len = len(prompt_tokens) - 1
    input_tokens = prompt_tokens + sample.tokens[:-1]
    target_tokens = [0] * observation_len + sample.tokens
    padded_logprobs = [0.0] * observation_len + sample.logprobs
    padded_advantages = (
        [0.0] * observation_len
        + [sample.advantage] * len(sample.tokens)
    )

    if not (
        len(input_tokens)
        == len(target_tokens)
        == len(padded_logprobs)
        == len(padded_advantages)
    ):
        raise ValueError("GRPO datum fields must have the same token length")

    return 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(padded_logprobs, dtype=np.float32),
            "advantages": np.asarray(padded_advantages, dtype=np.float32),
        },
    )

prompt 区间只提供上下文,因此 target_tokenslogprobsadvantages 使用 0 占位。completion 区间放入真实目标 token、旧策略 logprob 和整条回答共享的组内 advantage。

以 3 个 prompt token 和 2 个 completion token 为例:

prompt       = [x1, x2, x3]
completion   = [y1, y2]

model_input  = [x1, x2, x3, y1]
target       = [ 0,  0, y1, y2]
old_logprob  = [ 0,  0, l1, l2]
advantage    = [ 0,  0,  A,  A]

最后一个 prompt token x3 对应的位置开始预测 y1,因此 prompt mask 的长度是 len(prompt_tokens) - 1

(5)选择 importance sampling 或 PPO

importance_samplingppo 复用同一批 GRPO Datum,只需切换 loss_fn。两个 loss 都由 PyTRIO 内置实现:

fwd_bwd_future = training_client.forward_backward(
    datums,
    loss_fn=config.loss_fn,
)

选择 importance_sampling 时,训练直接使用新旧策略概率比乘以 advantage。选择 ppo 时,PyTRIO 在相同数据上应用 PPO clipping。GRPO 的 rollout、reward 与组内 advantage 构造保持不变。

(6)串联一轮同步训练

main() 先创建训练客户端、采样参数与优化器:

service_client = trio.ServiceClient()
training_client = service_client.create_lora_training_client(
    base_model=config.base_model,
    rank=config.lora_rank,
)
tokenizer = training_client.get_tokenizer()

sampling_params = trio.SamplingParams(
    max_tokens=config.max_tokens,
    temperature=config.temperature,
    top_p=config.top_p,
    stop=get_stop_sequences(tokenizer),
)
adam_params = trio.AdamParams(
    learning_rate=config.learning_rate,
    beta1=config.beta1,
    beta2=config.beta2,
)

每个 step 都先从当前训练权重创建新的 sampler,然后按题目顺序完成 rollout:

for step in range(effective_steps):
    batch_rows = pick_batch(
        train_data,
        step,
        config.batch_size,
        config.all_data,
    )
    sampling_client = (
        training_client.save_weights_and_get_sampling_client()
    )

    datums: list[trio.Datum] = []
    prompt_mean_rewards: list[float] = []
    rollout_lengths: list[int] = []
    n_degenerate = 0

    for row in tqdm(
        batch_rows,
        desc=f"GRPO step {step}",
        unit="prompt",
    ):
        prompt_tokens = build_prompt(tokenizer, row["question"])
        ground_truth = extract_gsm8k_answer(row["answer"])
        rollout_samples = run_rollout_group(
            sampling_client=sampling_client,
            tokenizer=tokenizer,
            prompt_tokens=prompt_tokens,
            ground_truth=ground_truth,
            sampling_params=sampling_params,
            group_size=config.group_size,
        )

        rewards = [sample.reward for sample in rollout_samples]
        prompt_mean_rewards.append(sum(rewards) / len(rewards))
        rollout_lengths.extend(
            len(sample.tokens) for sample in rollout_samples
        )

        if all(
            sample.advantage == 0.0
            for sample in rollout_samples
        ):
            n_degenerate += 1
            continue

        for sample in rollout_samples:
            datums.append(
                build_grpo_datum(prompt_tokens, sample)
            )

采样器的刷新位置体现了 on-policy 约束:第 $k$ 步更新完成后,第 $k+1$ 步用新权重重新生成轨迹。同步版在提交前向反向与优化任务后显式等待结果:

optim_future = training_client.optim_step(adam_params)
fwd_bwd_result = fwd_bwd_future.result()
optim_future.result()
loss_metrics = dict(fwd_bwd_result.metrics)

脚本最后统计 batch 平均 reward、退化组比例、平均生成长度、有效训练 token 数与 loss_mean,并保存可用于采样的 LoRA 权重:

final_weights = training_client.save_weights_for_sampler(
    name=run_name
).result()
print(
    f"Saved weights name: {run_name}, "
    f"path: {final_weights.path}"
)

8.1.6 运行同步版与异步版

先登录 TRIO,再从仓库根目录运行同步版的一步试验:

trio login

python docs/chapter8/grpo/01-demo-sync.py \
    --steps 1 \
    --batch-size 1 \
    --group-size 4 \
    --max-tokens 512 \
    --loss-fn importance_sampling \
    --swanlab-mode disabled

异步版使用完全相同的训练参数:

python docs/chapter8/grpo/02-demo-async.py \
    --steps 1 \
    --batch-size 4 \
    --group-size 4 \
    --max-tokens 512 \
    --loss-fn importance_sampling \
    --swanlab-mode disabled

异步版的算法数据没有变化,执行方式发生了以下变化:

训练环节同步版异步版
单题 rolloutsample(...).result()await sample_async(...)
batch 内多题依次处理asyncio.gather() 并发处理
刷新 sampler同步返回await save_weights_and_get_sampling_client_async()
前向反向forward_backward().result()forward_backward_async() 后再次 await future
优化器optim_step().result()optim_step_async() 后再次 await future
程序入口main(config)asyncio.run(main(config))

正文使用同步版建立算法直觉。异步版适合 batch 中含有多个独立 prompt 的正式实验,它通过并发远程请求减少串行等待时间。

训练时应优先观察代码真实记录的指标:

指标含义异常现象
rewardbatch 内各 prompt 平均正确率的均值长期不升时检查 reward、数据难度和学习率
frac_degenerate全对或全错 group 的比例过高时有效相对优势不足
rollout/avg_gen_lencompletion 平均 token 数接近 max_tokens 时检查截断
train_tokensadvantage 非零的 completion token 数为 0 时当前 step 没有训练信号
loss_mean当前策略更新的主 loss剧烈波动时检查学习率与概率比

至此,我们得到了本章的第一条训练主线:当前策略分组采样,规则环境给出结果奖励,组内比较产生 advantage,PyTRIO 完成策略更新。接下来的 OPD 会保留同样的 on-policy 数据流,同时把训练信号换成 Teacher 对每个 token 的反馈。

8.2 On-Policy Distillation

代码来源: 本节依据作者开源仓库 agentic-rl-lab/02-opd/general-opd 中的 01-demo-sync.py02-demo-async.py 整理。正文逐段讲解同步版 01-demo-sync.py,异步版作为完整配套实现保留。

大模型能力蒸馏通常由一个较强的 Teacher 和一个较小的 Student 组成。传统知识蒸馏使用预先收集好的数据,让 Student 拟合 Teacher 在这些固定样本上的概率分布。随着 Student 在训练中不断变化,固定数据与 Student 真正会访问的状态之间可能出现偏差。

On-Policy Distillation(OPD)让 Student 使用当前策略生成回答,再由 Teacher 对 Student 已经走过的轨迹逐 token 打分。这样,每一轮监督都位于 Student 当前的状态分布上。Teacher 无需重新生成一套答案,也无需参与反向传播,它只为 Student 的动作提供稠密概率反馈。

OPD 近年逐渐成为大模型训练中的重要方向。它可以用于推理能力迁移、能力恢复以及持续学习,还能够与强化学习的数据流自然结合:Student rollout 对应 on-policy 轨迹,Teacher logprob 对应逐 token 训练信号,Student 更新后再生成下一批轨迹。

8.2.1 从离线蒸馏到 On-Policy Distillation

设 Teacher 为 $\pi_T$,Student 为 $\pi_\theta$。离线蒸馏先固定一个数据集 $\mathcal D$,然后在其中的状态上最小化分布差异:

$$\mathcal L_{\mathrm{offline}} = \mathbb E_{s\sim\mathcal D} \left[D\bigl(\pi_T(\cdot\mid s),\pi_\theta(\cdot\mid s)\bigr)\right].$$

这种做法实现简单,效果取决于数据集能否覆盖 Student 推理时遇到的状态。假设固定数据中大多是 Teacher 的高质量轨迹,而 Student 在真实生成时较早出现了一个错误 token,它随后进入的状态可能从未出现在蒸馏数据中,Teacher 的指导便无法覆盖这段分布外轨迹。

OPD 把状态采样改为当前 Student:

$$y\sim\pi_{\theta_{\mathrm{old}}}(\cdot\mid x),$$

然后让 Teacher 计算同一条 Student completion 的条件概率:

$$\log\pi_T(y_t\mid x,y_{\lt t}).$$

一次训练迭代可以分为三步:

  1. Student rollout:Student 使用当前权重生成回答;
  2. Teacher feedback:Teacher 在完全相同的 prompt 与 completion 上计算 logprob;
  3. Student update:根据两者的概率差构造逐 token 信号并更新 Student。
OPD 在 Student 自身轨迹上获得 Teacher 反馈

图 8.3 OPD 在 Student 自身轨迹上获得 Teacher 反馈

图 8.3 中,Teacher 只进行前向计算。Student 的采样动作也不参与反向传播,梯度从构造好的训练目标回到当前 Student 参数。这种停止梯度的边界使 OPD 可以直接复用现有 rollout 与策略优化基础设施。

8.2.2 Reverse KL 与逐 token 训练信号

OPD 可以选择不同的分布散度。本章使用 reverse KL:

$$D_{\mathrm{KL}}\left(\pi_\theta,|,\pi_T\right) = \mathbb E_{y\sim\pi_\theta} \left[\log\pi_\theta(y\mid x)-\log\pi_T(y\mid x)\right].$$

对于 Student 实际采样到的第 $t$ 个 token,可以得到一个 Monte Carlo 估计:

$$d_t = \log\pi_{\theta_{\mathrm{old}}}(y_t\mid x,y_{\lt t})-\log\pi_T(y_t\mid x,y_{\lt t}).$$

当 $d_t>0$ 时,Student 对这个 token 的偏好高于 Teacher;当 $d_t<0$ 时,Teacher 给出的概率更高。为了最小化 reverse KL,本章将训练 advantage 写为:

$$A_t=-\beta d_t,$$

其中, $\beta$ 控制蒸馏强度。这样,每个 completion token 都能获得独立的训练信号,信号密度远高于只在回答结束时给出一个 0/1 reward。

reverse KL 具有较强的 mode-seeking 特性。Student 会优先集中到 Teacher 已经赋予较高概率的区域,适合将 Teacher 的高置信行为压缩到小模型中。与此同时,Teacher 与 Student 的能力差距、推理风格和 tokenizer 兼容性都会影响蒸馏效果。更大的 Teacher 只是候选条件,Teacher 还需要在 Student 当前轨迹上提供有价值且可学习的新信号。

表 8.2 对比了 GRPO 与本章 OPD 示例的数据来源。

项目GRPOOPD
轨迹来源当前策略当前 Student
反馈来源环境或规则判题器Teacher 概率分布
信号粒度常见为整条轨迹 reward逐 completion token
是否需要标准答案通常需要可验证结果prompt-only 数据即可
是否需要额外模型可省去 Value Model需要 Teacher 前向计算
主要风险reward hacking、退化组Teacher 不兼容、能力差距不合适

8.2.3 Teacher 必须评价同一条 Student 轨迹

OPD 的关键数据约束,是 Teacher 与 Student 必须在相同前缀下评价相同目标 token。Teacher 自己生成一条回答,再让 Student 学习这条回答,会退回到常见的合成数据蒸馏流程。

在 PyTRIO 中,可以先把 prompt 和 Student completion 拼接起来,再调用 Teacher 的 compute_logprobs()

def completion_teacher_logprobs(
    teacher_client,
    prompt_ids: list[int],
    completion_ids: list[int],
):
    all_ids = prompt_ids + completion_ids
    all_logprobs = teacher_client.compute_logprobs(
        trio.ModelInput.from_ints(all_ids)
    ).result()

    start = len(prompt_ids)
    end = start + len(completion_ids)
    return [float(x) for x in all_logprobs[start:end]]

不同 SDK 对 logprob 的位置定义可能存在一位偏移,因此不能仅凭数组长度猜测切片。最稳妥的方法是构造一个很短的已知序列,确认返回值中每个位置究竟表示“当前 token 的 logprob”还是“下一个 token 的 logprob”,再完成 prompt 与 completion 的边界测试。本章代码已经按照 PyTRIO 0.2.6 的返回约定完成对齐。

Teacher 和 Student 还需要共享可兼容的 tokenizer。若同一组 token id 在两个模型中代表不同字符串,即使数组长度相同,概率差也没有正确语义。正式训练前应验证:

text = "A short tokenizer alignment test."
student_ids = student_tokenizer.encode(text, add_special_tokens=False)
teacher_ids = teacher_tokenizer.encode(text, add_special_tokens=False)

assert student_ids == teacher_ids

对于同系列但 chat template 不同的模型,应由 Student tokenizer 渲染一次完整 prompt,然后让 Teacher 直接评价对应 token 序列。不要分别套用两份 chat template,否则两边看到的上下文已经发生变化。

8.2.4 从头实现同步版 OPD

OPD 配套目录同样提供同步和异步两个独立脚本:

文件执行方式用途
01-demo-sync.py同步正文逐段讲解的教学主线
02-demo-async.py异步并发处理 Student rollout 与 Teacher logprob 请求

下面按照 01-demo-sync.py 的代码顺序,实现从 DeepMath prompt 到 Student 更新的一轮完整 OPD。

(1)准备参数与 DeepMath 数据

脚本默认把数据下载到当前代码目录下的 datasets/DeepMath-103K

SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_DATA_DIR = SCRIPT_DIR / "datasets" / "DeepMath-103K"
DEEPMATH_SHARDS = 10

parse_args() 的参数可以分成四组:

  • 数据参数:dataset_repodataset_dirnum_shardssample_size
  • 模型参数:Student 基座、LoRA rank、Teacher 基座或 Teacher 权重路径;
  • 训练参数:step、batch、group、采样长度、KL 系数与优化器参数;
  • 记录参数:SwanLab 项目、实验名和日志模式。

DeepMath-103K 在 ModelScope 上包含 10 个 parquet 分片。shard_name() 生成远端文件名,modelscope_file_url() 生成下载地址,download_if_needed() 先写入临时文件再原子替换,避免中断时留下不完整分片。

数据加载函数把本地 parquet 交给 Hugging Face Datasets,只保留训练需要的 question 字段:

def load_deepmath(args: argparse.Namespace):
    shard_paths = []

    for index in range(args.num_shards):
        remote_path = shard_name(index)
        local_path = args.dataset_dir / remote_path
        download_if_needed(
            modelscope_file_url(
                args.dataset_repo,
                args.dataset_revision,
                remote_path,
            ),
            local_path,
            args.force_download,
        )
        shard_paths.append(str(local_path))

    dataset = load_dataset(
        "parquet",
        data_files=shard_paths,
        split="train",
        cache_dir=str(
            args.dataset_dir / ".datasets_cache"
        ),
    )

    if "question" not in dataset.column_names:
        raise ValueError(
            "DeepMath must contain 'question', "
            f"got {dataset.column_names}"
        )

    dataset = dataset.shuffle(seed=args.seed)
    if args.sample_size > 0:
        dataset = dataset.select(
            range(min(args.sample_size, len(dataset)))
        )
    return dataset

--num-shards 1 适合第一次联调,--num-shards 10 会覆盖完整数据集。--sample-size 控制加载后实际参与 prompt pool 的样本数。

(2)把问题渲染成 Student prompt

OPD 使用 prompt-only 数据,不读取 DeepMath 的参考推理过程。build_prompt() 把问题和输出要求交给 Student tokenizer:

def build_prompt(
    tokenizer,
    question: str,
    suffix: str,
    enable_thinking: bool,
) -> list[int]:
    content = (
        question.strip()
        if not suffix
        else f"{question.strip()}\n\n{suffix}"
    )
    messages = [{"role": "user", "content": content}]
    prompt = tokenizer.apply_chat_template(
        messages,
        tokenize=False,
        add_generation_prompt=True,
        enable_thinking=enable_thinking,
    )
    return tokenizer.encode(
        prompt,
        add_special_tokens=False,
    )

prompt 只渲染一次。Teacher 随后直接评价同一组 token,从而保持两边上下文完全一致。

(3)创建 Student 与 Teacher

Student 使用 LoRA training client,Teacher 使用只做前向计算的 sampling client:

service_client = trio.ServiceClient()

training_client = (
    service_client.create_lora_training_client(
        base_model=args.base_model,
        rank=args.lora_rank,
        seed=args.seed,
    )
)
tokenizer = training_client.get_tokenizer()

teacher_client = service_client.create_sampling_client(
    base_model=(
        args.teacher_base_model
        or args.base_model
    ),
    model_path=args.teacher_model_path,
)

teacher_model_path 可以指向已经保存的 sampler 权重。Teacher 不调用 forward_backward()optim_step(),它的参数在 OPD 过程中保持固定。

(4)让 Teacher 评价 Student 的同一条轨迹

Student rollout 返回 completion token 与旧策略 logprob。Teacher 需要看到 prompt_ids + completion_ids,再从返回结果中切出 completion 区间:

def completion_teacher_logprobs(
    teacher_client,
    prompt_ids: list[int],
    completion_ids: list[int],
):
    all_ids = prompt_ids + completion_ids
    all_logprobs = teacher_client.compute_logprobs(
        trio.ModelInput.from_ints(all_ids)
    ).result()

    completion_logprobs = all_logprobs[
        len(prompt_ids):
    ]

    if (
        len(completion_logprobs)
        != len(completion_ids)
        or any(
            value is None
            for value in completion_logprobs
        )
    ):
        raise ValueError(
            "Invalid teacher logprobs "
            "for completion tokens"
        )

    return [
        float(value)
        for value in completion_logprobs
    ]

长度检查与 None 检查用于保护 token 对齐。通过检查后,第 $t$ 个 Student completion token、Student 旧 logprob 和 Teacher logprob 才能组成同一个训练位置。

(5)构造逐 token OPD Datum

OPD 与 GRPO 使用相同的 importance_sampling schema,差异位于 advantage。GRPO 为整条 completion 复制一个组内 advantage,OPD 为每个 completion token 写入独立的 reverse-KL advantage:

def build_opd_datum(
    prompt_ids: list[int],
    completion_ids: list[int],
    old_logprobs,
    advantages,
):
    prompt_loss_len = len(prompt_ids) - 1

    input_ids = (
        prompt_ids
        + completion_ids[:-1]
    )
    target_ids = (
        [0] * prompt_loss_len
        + completion_ids
    )
    padded_logprobs = (
        [0.0] * prompt_loss_len
        + list(old_logprobs)
    )
    padded_advantages = (
        [0.0] * prompt_loss_len
        + list(advantages)
    )

    if not (
        len(input_ids)
        == len(target_ids)
        == len(padded_logprobs)
        == len(padded_advantages)
    ):
        raise ValueError(
            "OPD datum fields must "
            "have the same length"
        )

    return trio.Datum(
        model_input=trio.ModelInput.from_ints(
            input_ids
        ),
        loss_fn_inputs={
            "target_tokens": np.asarray(
                target_ids,
                dtype=np.int64,
            ),
            "logprobs": np.asarray(
                padded_logprobs,
                dtype=np.float32,
            ),
            "advantages": np.asarray(
                padded_advantages,
                dtype=np.float32,
            ),
        },
    )

右移规则与 GRPO 完全一致。prompt 区间的三个 loss 输入继续使用 0 占位,completion 区间一一放入 Student token、Student old logprob 和该 token 的 OPD advantage。

(6)生成 Student 轨迹并计算 reverse KL

每个 step 开始时,代码按 sampler_refresh_steps 刷新 Student sampler。默认值为 1,因此每一步都使用最新 Student 权重:

if (
    student_sampler is None
    or step % args.sampler_refresh_steps == 0
):
    student_sampler = (
        training_client
        .save_weights_and_get_sampling_client()
    )

同步版依次处理当前 batch 的问题。对每个 prompt,Student 先生成 group_size 条轨迹,Teacher 再对每条非空 completion 打分:

result = student_sampler.sample(
    prompt=trio.ModelInput.from_ints(
        prompt_ids
    ),
    num_samples=args.group_size,
    sampling_params=sampling_params,
    return_text=False,
).result()

for sequence in result.sequences:
    completion_ids = sequence.tokens
    if not completion_ids:
        continue

    student_lps = [
        float(value)
        for value in sequence.logprobs
    ]
    teacher_lps = completion_teacher_logprobs(
        teacher_client,
        prompt_ids,
        completion_ids,
    )

    reverse_kl = (
        np.asarray(student_lps)
        - np.asarray(teacher_lps)
    )
    advantages = (
        -args.kl_penalty_coef
        * reverse_kl
    )

    datums.append(
        build_opd_datum(
            prompt_ids,
            completion_ids,
            student_lps,
            advantages,
        )
    )
    reverse_kls.extend(reverse_kl.tolist())
    completion_token_counts.append(
        len(completion_ids)
    )

group_size 在这里扩大同一 prompt 下的 Student 状态覆盖范围。每条轨迹独立得到 Teacher 信号,不进行 GRPO 式的组内中心化。

(7)更新 Student、记录指标并保存权重

收集完整个 batch 后,代码执行一次前向反向和一次优化器更新:

if not datums:
    raise RuntimeError(
        "No OPD datums were built"
    )

fwd_bwd = training_client.forward_backward(
    datums,
    loss_fn="importance_sampling",
)
optim = training_client.optim_step(adam)

fwd_bwd_result = fwd_bwd.result()
optim.result()

一次 OPD step 的 tokens/s 使用 completion token 总数除以整个 step 的耗时。这个耗时同时包含 Student 采样、Teacher logprob、Student 前向反向、优化器更新以及必要的 sampler 刷新:

completion_tokens_total = int(
    sum(completion_token_counts)
)
metrics = {
    "data/datums": len(datums),
    "data/completion_tokens_mean": float(
        np.mean(completion_token_counts)
    ),
    "data/completion_tokens_total": (
        completion_tokens_total
    ),
    "data/completion_tokens_per_second": (
        completion_tokens_total
        / step_elapsed_time
    ),
    "opd/reverse_kl_mean": float(
        np.mean(reverse_kls)
    ),
    "opd/reverse_kl_std": float(
        np.std(reverse_kls)
    ),
    "train/learning_rate": args.learning_rate,
    "time/step_elapsed_time": step_elapsed_time,
}

训练结束后保存 Student 的 sampler 权重:

save_result = (
    training_client
    .save_weights_for_sampler(
        args.save_weights_name
    )
    .result()
)
print(f"Saved weights: {save_result.path}")

8.2.5 运行同步版与异步版

同步版的一步试验会自动下载一个 DeepMath 分片,并从中抽取最多 20 个 prompt:

python docs/chapter8/opd/01-demo-sync.py \
    --steps 1 \
    --batch-size 1 \
    --group-size 1 \
    --max-tokens 512 \
    --sample-size 20 \
    --num-shards 1 \
    --swanlab-mode disabled

异步版使用相同的算法与参数:

python docs/chapter8/opd/02-demo-async.py \
    --steps 1 \
    --batch-size 4 \
    --group-size 2 \
    --max-tokens 512 \
    --sample-size 20 \
    --num-shards 1 \
    --swanlab-mode disabled

异步 OPD 同时利用两层并发:

训练环节同步版异步版
创建 Student / Teacher同步 client 接口create_lora_training_client_async()create_sampling_client_async()
Student rollout逐 prompt 调用 sample().result()batch 内并发 sample_async()
Teacher 打分逐 completion 调用 compute_logprobs().result()同一 prompt 内并发 compute_logprobs_async()
汇总 batch顺序追加 Datum两层 asyncio.gather() 后统一汇总
Student 更新同步 future + .result()async API 提交后再次 await future
程序入口main(args)asyncio.run(main(args))

正文只展开同步版,因为它能直接呈现“Student 采样 → Teacher 打分 → reverse KL → Student 更新”的因果顺序。异步版保留同样的 token 对齐、advantage 和优化目标。

示例默认使用 Qwen/Qwen3.5-4B 作为 Student,使用 Qwen/Qwen3.6-27B 作为 Teacher。实际可用模型应以当前 PyTRIO 服务返回的支持列表为准。训练时应同步观察:

  • opd/reverse_kl_mean:Student 采样 token 与 Teacher 概率的平均差异;
  • opd/reverse_kl_std:不同 token 蒸馏信号的离散程度;
  • data/completion_tokens_total:本 step 完成 Teacher 打分的 token 总量;
  • data/completion_tokens_per_second:整步 OPD 的有效 completion token 吞吐;
  • trainer/*:PyTRIO 返回的训练指标;
  • 独立任务评测与通用能力评测:确认蒸馏后的真实能力变化。

reverse KL 下降说明 Student 的采样分布正在接近 Teacher。任务正确率仍需通过固定的 held-out 评测确认。出现 loss 波动或能力下降时,应依次检查 token 对齐、Teacher 质量、KL 系数、学习率、sampler 刷新频率和 completion 截断比例。

GRPO 和 OPD 到这里都发生在“模型生成一段回答,外部模块提供训练信号”的单轮结构中。下一节的 Search-R1 会进一步改变轨迹形态:模型在完成回答前可以暂停生成、调用搜索环境、读取 observation,再继续决定下一步动作。

8.3 Search-R1

代码来源: 本节依据作者开源仓库 agentic-rl-lab/03-search-r1 中的完整多文件实现整理。正文按照 prepare_data.pyprotocol.pysearch.pyrollout.pyreward.pytrain.pyeval.pyanalyse.py 的真实调用链展开。

大语言模型内部存储了大量知识,但这些知识受训练数据时间范围、参数容量和长尾覆盖率限制。检索增强生成(Retrieval-Augmented Generation,RAG)通常由系统预先检索文档,再把结果交给模型回答。Search-R1 将搜索决策交给模型:模型可以判断何时搜索、搜索什么、是否需要根据新证据继续搜索,以及何时结束并给出答案。

这项变化使一次回答成为多轮轨迹:

$$\tau=(x,a_1,o_1,a_2,o_2,\ldots,a_T),$$

其中, $a_t$ 是模型生成的搜索调用或最终回答, $o_t$ 是搜索环境返回的 observation。模型只有在完成整条轨迹后才能获得结果奖励,GRPO 再把成功搜索路径与失败路径区分开。

配套实现不是一个孤立的训练脚本,而是一条由 9 个文件共同组成的执行链。阅读代码时,可以先抓住下面这条主线:

prepare_data.py / data.py
        ↓ 读取问题与参考答案
protocol.py
        ↓ 构造 prompt、解析 search 调用
search.py
        ↓ 执行检索并返回 observation
rollout.py
        ↓ 生成完整多轮轨迹、reward 与组内 advantage
train.py
        ↓ observation mask、Datum、importance_sampling、optim_step
eval.py / analyse.py
        ↓ 固定环境评测与结果汇总

真正被训练的是语言模型生成的 Assistant token。search.py 负责的搜索服务属于环境,protocol.py 负责把模型动作和环境 observation 接成连续轨迹,rollout.py 负责调度两者,train.py 最后才把轨迹转换成 PyTRIO 可以更新的 Datum。下面按照这条真实调用链展开代码。

8.3.1 从固定检索到自主搜索

Search-R1 的训练样本只包含问题和参考答案,不包含人工编写的搜索词,也不包含标准搜索轨迹。prepare_data.py 将 NQ 与 HotpotQA 的原始字段统一为下面的 JSONL 结构:

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

清洗逻辑只保留可用于结果奖励的字段:

def normalize_row(row: dict[str, Any]) -> dict[str, Any] | None:
    question = str(row.get("question") or "").strip()
    answers = extract_answers(row)
    if not question or not answers:
        return None
    return {
        "id": str(row.get("id") or ""),
        "question": question,
        "answers": answers,
        "data_source": str(
            row.get("data_source") or "unknown"
        ),
    }

data.py 再将每一行转换为 SearchExample

@dataclass(frozen=True)
class SearchExample:
    id: str
    question: str
    answers: list[str]
    data_source: str

这个数据边界很重要。搜索 query、搜索次数和证据使用方式都必须由当前策略在 rollout 中自己探索;如果训练集已经给出了固定检索轨迹,任务就更接近轨迹模仿,而不是本节讨论的 Search-R1。

固定 RAG 与 Search Agent 的差异可以从表 8.3 中看出。

环节固定 RAGSearch-R1
是否检索由外部流程决定模型根据当前状态决定
检索词预设或单次生成可根据 observation 多次改写
检索轮数通常固定由策略决定,受最大次数约束
训练对象回答模型或检索器搜索动作与最终回答的统一策略
奖励常见为监督答案 loss完整轨迹的结果奖励

例如,问题“《小王子》的作者出生在哪个国家?”需要先确定作者,再查询作者的出生地。模型可以形成下面的轨迹:

User: What country was the author of The Little Prince born in?

Assistant -> search("The Little Prince author")
Tool: Antoine de Saint-Exupéry ...

Assistant -> search("Antoine de Saint-Exupéry birthplace country")
Tool: He was born in Lyon, France ...

Assistant: Answer: France

第二次搜索词依赖第一次 observation,因此无法在 rollout 开始前一次性确定。训练代码需要维护环境状态,并在每个 assistant turn 后解析模型动作。

Search-R1 的多轮搜索与 observation mask

图 8.4 Search-R1 的多轮搜索与 observation mask

Search-R1 论文使用 <search>... </search><information>... </information> 等文本标签表示动作与观察。本章使用 Qwen3.5-4B 模型,采用模型原生 function calling 模板:

SEARCH_TOOL = {
    "type": "function",
    "function": {
        "name": "search",
        "description": "Search the web for evidence.",
        "parameters": {
            "type": "object",
            "properties": {
                "query": {"type": "string"},
            },
            "required": ["query"],
        },
    },
}

两种协议承载相同的强化学习语义:Assistant 输出可训练动作,搜索后端返回环境 observation。使用原生模板可以减少模型基础能力与工具格式之间的冲突,也能保留结构化的工具消息。

定义工具以后,还必须把模型输出严格解析为三种状态:合法搜索、最终回答或非法输出。protocol.py 的解析器要求一轮中只能出现一次完整的 search 调用,并且工具调用后不能再夹带最终答案:

def parse_assistant(text: str) -> ParsedAssistant:
    matches = list(TOOL_CALL_PATTERN.finditer(text))
    if not matches:
        kind = (
            "invalid" if "<tool_call>" in text else "answer"
        )
        return ParsedAssistant(
            kind=kind,
            content=text.strip(),
        )

    if len(matches) != 1 or text[matches[0].end():].strip():
        return ParsedAssistant(
            kind="invalid",
            content=text.strip(),
        )

    query = matches[0].group(1).strip()
    if not query or "<" in query or ">" in query:
        return ParsedAssistant(
            kind="invalid",
            content=text.strip(),
        )

    content = text[:matches[0].start()].strip()
    return ParsedAssistant(
        kind="tool",
        content=content,
        query=query,
    )

这里没有单独定义“停止”动作。只要输出中不存在合法工具调用,本轮就作为最终回答结束;之后由 reward.py 检查它是否符合 Answer: <short answer>。包含残缺 <tool_call> 的输出会被标记为 invalid,同样结束轨迹并得到格式惩罚。

8.3.2 搜索后端怎样变成环境 observation

search.py 提供 DeepSeek Search、Wikipedia 和知乎三个后端。rollout 不直接依赖任何 HTTP 响应格式,而只依赖统一的 SearchResult

@dataclass(frozen=True)
class SearchItem:
    title: str
    content: str
    source: str | None = None
    url: str | None = None


@dataclass(frozen=True)
class SearchResult:
    ok: bool
    items: list[SearchItem]
    latency: float
    status: int | None = None
    error: str | None = None

训练入口根据命令行参数创建后端,后面的状态机始终调用同一个 search(query) 接口:

def create_search_client(
    backend: str,
    env_path: str | Path | None = None,
    *,
    model: str = "deepseek-v4-flash",
    timeout: float | None = None,
) -> SearchClient:
    resolved_timeout = resolve_search_timeout(
        backend, timeout
    )
    if backend == "deepseek":
        return DeepSeekSearchClient.from_env(
            env_path,
            model=model,
            timeout=resolved_timeout,
        )
    if backend == "wikipedia":
        return WikipediaSearchClient(
            timeout=resolved_timeout
        )
    if backend == "zhihu":
        return ZhihuSearchClient.from_env(
            env_path,
            timeout=resolved_timeout,
        )
    raise ValueError(f"不支持的搜索后端: {backend}")

以无需密钥的 Wikipedia 后端为例,真正的检索发生在 _request() 中。一次 API 请求同时完成全文搜索,并取回前三个页面的标题、正文摘要和 URL:

def _request(
    self,
    query: str,
    started: float,
) -> SearchResult:
    params = urllib.parse.urlencode(
        {
            "action": "query",
            "format": "json",
            "formatversion": 2,
            "generator": "search",
            "gsrsearch": query,
            "gsrlimit": 3,
            "prop": "extracts|info",
            "explaintext": 1,
            "exintro": 1,
            "exchars": 1200,
            "inprop": "url",
            "redirects": 1,
            "utf8": 1,
        }
    )
    request = urllib.request.Request(
        f"{WIKIPEDIA_SEARCH_ENDPOINT}?{params}",
        headers={
            "Accept": "application/json",
            "Accept-Encoding": "gzip",
            "User-Agent": WIKIPEDIA_USER_AGENT,
        },
    )
    with urllib.request.urlopen(
        request,
        timeout=self.timeout,
    ) as response:
        body = response.read()
        if response.headers.get(
            "Content-Encoding"
        ) == "gzip":
            body = gzip.decompress(body)
        payload = json.loads(body.decode("utf-8"))
        pages = payload.get("query", {}).get(
            "pages", []
        )
        if not isinstance(pages, list):
            raise TypeError(
                "Wikipedia 搜索响应 pages 不是列表"
            )
        if any(
            not isinstance(page, dict)
            for page in pages
        ):
            raise TypeError(
                "Wikipedia 搜索响应 page 不是对象"
            )
        ordered_pages = sorted(
            pages,
            key=lambda page: int(
                page.get("index", 1_000_000)
            ),
        )
        items = [
            SearchItem(
                title=str(
                    page.get("title") or "Untitled"
                ).strip(),
                content=str(
                    page.get("extract") or ""
                ).strip(),
                source="Wikipedia",
                url=str(
                    page.get("fullurl") or ""
                ).strip(),
            )
            for page in ordered_pages
            if str(page.get("extract") or "").strip()
        ]
        return SearchResult(
            ok=True,
            items=items,
            latency=time.perf_counter() - started,
            status=response.status,
        )

外层 search() 还负责请求限速、timeout、429 与 5xx 的有限重试,并累计成功率和 latency。重试后仍失败时,它不会让整个 rollout 抛异常退出,而是返回 SearchResult(ok=False, error=...)rollout.py 会把错误文本也作为 observation 交给模型,因此模型可以继续回答或改写 query。

不同后端的结果最终都被格式化为可读证据:

[1] Title: ...
    Content: ...
    Source: ...
    URL: ...

只返回标题或链接通常不够,因为模型真正需要的是能够支持下一步推理的 contentmax_tool_response_tokensmax_trajectory_tokens 会在 observation 接入轨迹前再次限制证据长度。

8.3.3 多轮 rollout 状态机

单轮 GRPO 只需一次 sample()。Search-R1 每条轨迹都可能经历不同数量的搜索调用,训练器需要反复执行“生成—解析—调用环境—追加 observation”。

本章将一条轨迹保存为 Trajectory

@dataclass
class AssistantTurn:
    prompt_tokens: list[int]
    completion_tokens: list[int]
    logprobs: list[float]
    text: str


@dataclass
class Trajectory:
    example: SearchExample
    group_index: int
    messages: list[dict[str, Any]]
    next_prompt_tokens: list[int] | None = None
    question_index: int = 0
    turns: list[AssistantTurn] = field(default_factory=list)
    search_calls: int = 0
    final_text: str = ""
    reward: float = -0.1
    advantage: float = 0.0
    valid_format: bool = False
    exact_match: bool = False
    done: bool = False

其中,每个 AssistantTurn 记录本轮采样前的完整 prompt、模型生成 token 和对应旧策略 logprob。搜索结果存在 messages 和下一轮 prompt 中,但它没有 rollout logprob,因为这些 token 来自环境。

真正向 PyTRIO 发请求的是 sample_requests_async()。同一轮中,不同问题或已经分叉的不同轨迹可以并发采样:

async def sample_requests_async(
    sampling_client: Any,
    requests: list[SampleRequest],
    config: RolloutConfig,
    tokenizer: Any,
) -> list[Any]:
    tasks = []
    for request in requests:
        params = trio.SamplingParams(
            max_tokens=request.max_tokens,
            seed=request.seed,
            stop=stop_sequences(tokenizer),
            temperature=config.temperature,
            top_p=config.top_p,
        )
        tasks.append(
            sampling_client.sample_async(
                prompt=trio.ModelInput.from_ints(
                    request.prompt_tokens
                ),
                num_samples=request.num_samples,
                sampling_params=params,
                return_text=True,
            )
        )
    return list(await asyncio.gather(*tasks))

采样返回后,consume_assistant() 先保存原始 token 与 old logprob,再调用上一节的协议解析器。合法搜索会变成 PendingSearch;最终回答、非法格式或预算耗尽都会结束当前轨迹:

def consume_assistant(
    trajectory: Trajectory,
    prompt_tokens: list[int],
    sequence: Any,
    tokenizer: Any,
    config: RolloutConfig,
) -> PendingSearch | None:
    tokens, logprobs, text = read_sequence(
        sequence, tokenizer
    )
    trajectory.turns.append(
        AssistantTurn(
            prompt_tokens,
            tokens,
            logprobs,
            text,
        )
    )
    parsed = parse_assistant(text)

    can_search = (
        parsed.kind == "tool"
        and trajectory.search_calls
        < config.max_search_calls
        and len(trajectory.turns)
        < config.max_assistant_turns
    )
    if not can_search:
        trajectory.messages.append(
            {"role": "assistant", "content": text}
        )
        trajectory.final_text = text
        trajectory.done = True
        return None

    call_id = (
        f"search-{trajectory.question_index}-"
        f"{trajectory.group_index}-"
        f"{trajectory.search_calls + 1}"
    )
    messages_before_assistant = list(
        trajectory.messages
    )
    trajectory.messages.append(
        {"role": "assistant", "content": text}
    )
    return PendingSearch(
        trajectory=trajectory,
        messages_before_assistant=(
            messages_before_assistant
        ),
        assistant_text=text,
        prompt_tokens=prompt_tokens,
        completion_tokens=tokens,
        call_id=call_id,
        query=parsed.query or "",
    )

resolve_searches() 使用线程池并发调用不同轨迹的搜索后端,再按原顺序把结果和待处理轨迹配对:

def resolve_searches(
    pending_searches: list[PendingSearch],
    search_client: SearchClient,
    tokenizer: Any,
    config: RolloutConfig,
) -> int:
    if not pending_searches:
        return 0
    workers = min(
        config.search_concurrency,
        len(pending_searches),
    )
    with ThreadPoolExecutor(
        max_workers=workers
    ) as pool:
        results = list(
            pool.map(
                search_client.search,
                [
                    pending.query
                    for pending in pending_searches
                ],
            )
        )
    return sum(
        finish_search(
            pending,
            result,
            tokenizer,
            config,
        )
        for pending, result in zip(
            pending_searches,
            results,
            strict=True,
        )
    )

每个结果随后由 finish_search() 写回原来的轨迹:

def finish_search(
    pending: PendingSearch,
    result: SearchResult,
    tokenizer: Any,
    config: RolloutConfig,
) -> bool:
    fitted = fit_tool_content(
        tokenizer,
        pending.messages_before_assistant,
        pending.assistant_text,
        pending.prompt_tokens,
        pending.completion_tokens,
        pending.call_id,
        result,
        config,
    )
    if fitted is None:
        pending.trajectory.final_text = (
            pending.assistant_text
        )
        pending.trajectory.done = True
        return True

    content, next_prompt_tokens = fitted
    pending.trajectory.messages.append(
        tool_message(pending.call_id, content)
    )
    pending.trajectory.next_prompt_tokens = (
        next_prompt_tokens
    )
    pending.trajectory.search_calls += 1
    return False

fit_tool_content() 会逐条加入搜索证据。只要 observation 超过单次工具预算或完整轨迹预算,就停止加入后续证据;如果连一条结果都放不下,轨迹会直接结束。这样,搜索服务返回多少文本和模型实际能看到多少文本是两个明确分开的边界。

一组 GRPO 轨迹在第一轮共享完全相同的 prompt。代码用一次请求生成 group_size 个分支:

first_requests: list[SampleRequest] = []
for index, trajectory in enumerate(roots):
    request = make_request(
        tokenizer,
        trajectory,
        index,
        config.group_size,
        config.seed + index,
        config,
    )
    if request:
        first_requests.append(request)

采样响应中的每条 sequence 都从根轨迹深拷贝出一个独立分支,随后再分别解析其搜索动作:

responses = asyncio.run(
    sample_requests_async(
        sampling_client,
        first_requests,
        config,
        tokenizer,
    )
)
pending_searches: list[PendingSearch] = []
for request, response in zip(
    first_requests,
    responses,
    strict=True,
):
    root = roots[request.trajectory_index]
    if len(response.sequences) != config.group_size:
        raise ValueError(
            "首轮采样数量与 group_size 不一致"
        )

    for group_index, sequence in enumerate(
        response.sequences
    ):
        branch = copy.deepcopy(root)
        branch.group_index = group_index
        pending = consume_assistant(
            branch,
            request.prompt_tokens,
            sequence,
            tokenizer,
            config,
        )
        trajectories.append(branch)
        if pending is not None:
            pending_searches.append(pending)

resolve_searches(
    pending_searches,
    search_client,
    tokenizer,
    config,
)

不同分支执行搜索后,observation 与历史动作已经不同。后续轮次对每条未结束轨迹分别设置 num_samples=1

while any(not item.done for item in trajectories):
    requests: list[SampleRequest] = []
    for index, trajectory in enumerate(trajectories):
        if trajectory.done:
            continue
        request = make_request(
            tokenizer,
            trajectory,
            index,
            1,
            (
                config.seed
                + index
                + len(trajectory.turns) * 10_000
            ),
            config,
        )
        if request:
            requests.append(request)

    if not requests:
        break
    responses = asyncio.run(
        sample_requests_async(
            sampling_client,
            requests,
            config,
            tokenizer,
        )
    )

    pending_searches = []
    for request, response in zip(
        requests, responses, strict=True
    ):
        if len(response.sequences) != 1:
            raise ValueError(
                "后续轮次必须只采样一个分支"
            )
        trajectory = trajectories[
            request.trajectory_index
        ]
        pending = consume_assistant(
            trajectory,
            request.prompt_tokens,
            response.sequences[0],
            tokenizer,
            config,
        )
        if pending is not None:
            pending_searches.append(pending)

    resolve_searches(
        pending_searches,
        search_client,
        tokenizer,
        config,
    )

这种“共享根节点、随后分叉”的实现同时满足两个条件:同题轨迹能够进行 GRPO 组内比较,每条分支又能根据自己的搜索结果继续行动。多个 PyTRIO 采样请求和搜索请求都可以并发执行,从而减少多轮环境交互带来的等待时间。

状态机还必须设置明确边界:

  • max_search_calls:一条轨迹最多调用多少次搜索;
  • max_assistant_turns:最多生成多少轮 Assistant 消息;
  • max_assistant_tokens:单轮动作的 token 上限;
  • max_tool_response_tokens:单次 observation 的 token 上限;
  • max_trajectory_tokens:整条上下文的总 token 上限;
  • 搜索 timeout 和 concurrency:限制外部服务等待时间与并发压力。

这些限制既控制训练成本,也定义了 Agent 可以访问的环境范围。评测时必须保持相同配置,否则更高的搜索预算本身就可能提高答案正确率。

8.3.4 工具协议与 token 连续性

多轮训练最容易出现的问题,是重新渲染 chat template 后的 token 与真实采样 token 不一致。tokenizer 可能清理空格、规范工具调用文本,或自动补充轮次结束符。如果训练阶段重新编码模型文本,旧 logprob 将无法和目标 token 一一对应。

本章遵循两个原则:

  1. Assistant 动作始终保留 sampler 返回的原始 token;
  2. 只通过 chat template 计算新加入的结束符与 tool observation token。

build_next_prompt() 的完整职责,是从 chat template 中只提取“Assistant 结束符 + 新 observation”这一段增量 token,再接到 sampler 返回的真实 completion 后面:

def build_next_prompt(
    tokenizer: Any,
    messages_before_assistant: list[dict[str, Any]],
    assistant_text: str,
    previous_prompt_tokens: list[int],
    completion_tokens: list[int],
    next_tool_message: dict[str, Any],
) -> list[int]:
    canonical_prompt = build_prompt(
        tokenizer,
        messages_before_assistant,
    )

    empty_assistant_end = _render_chat(
        tokenizer,
        [
            *messages_before_assistant,
            {"role": "assistant", "content": ""},
        ],
        add_generation_prompt=False,
    )
    if empty_assistant_end[
        :len(canonical_prompt)
    ] != canonical_prompt:
        raise ValueError(
            "chat template 无法从空 assistant 提取结束边界"
        )
    assistant_closing_tokens = empty_assistant_end[
        len(canonical_prompt):
    ]

    assistant_message = {
        "role": "assistant",
        "content": assistant_text,
    }
    messages_with_assistant = [
        *messages_before_assistant,
        assistant_message,
    ]
    canonical_assistant_end = _render_chat(
        tokenizer,
        messages_with_assistant,
        add_generation_prompt=False,
    )
    canonical_next_prompt = build_prompt(
        tokenizer,
        [*messages_with_assistant, next_tool_message],
    )
    if canonical_next_prompt[
        :len(canonical_assistant_end)
    ] != canonical_assistant_end:
        raise ValueError(
            "加入 tool observation 后 chat template 改写了历史消息"
        )

    observation_tokens = canonical_next_prompt[
        len(canonical_assistant_end):
    ]
    overlap = _suffix_prefix_overlap(
        completion_tokens,
        assistant_closing_tokens,
    )
    return [
        *previous_prompt_tokens,
        *completion_tokens,
        *assistant_closing_tokens[overlap:],
        *observation_tokens,
    ]

overlap 用来处理 sampler 已经返回部分 Assistant 结束符的情况,避免重复补 token。messages 负责协议记录与调试,next_prompt_tokens 才是下一轮真正提交给 sampler 的模型输入;两者不能混为一谈。

到下一轮训练数据构造时,还会再次验证:

if turn.prompt_tokens[:len(full_tokens)] != full_tokens:
    raise ValueError(
        "下一轮 prompt 不是已有轨迹的前缀扩展"
    )

这两层校验能够尽早发现模板改写、结束符重复和 observation 拼接错误。对于 Agentic RL,token 连续性属于算法正确性的一部分;轨迹在文本层面看起来合理,仍可能因 token 边界错位而产生错误梯度。

8.3.5 Reward、组内优势与 observation mask

本章使用 NQ 和 HotpotQA 的短答案任务。模型最终需要输出唯一一行:

Answer: <short answer>

规则奖励分为三档:

$$R(\tau)=\begin{cases}1, & \text{格式正确且答案精确匹配},\ 0, & \text{格式正确但答案错误},\ -0.1, & \text{最终格式无效}.\end{cases}$$

reward.py 先抽取唯一一行最终答案,再统一大小写、标点、英文冠词和空格:

def normalize_answer(text: str) -> str:
    lowered = text.lower()
    without_punctuation = "".join(
        char
        for char in lowered
        if not unicodedata.category(char).startswith("P")
    )
    without_articles = ARTICLE_PATTERN.sub(
        " ", without_punctuation
    )
    return " ".join(without_articles.split())


def extract_answer(text: str) -> str | None:
    matches = ANSWER_PATTERN.findall(text)
    if len(matches) != 1:
        return None
    answer = matches[0].strip()
    return answer or None


def score_answer(
    text: str,
    references: list[str],
) -> RewardResult:
    answer = extract_answer(text)
    if answer is None:
        return RewardResult(
            -0.1, False, False, None
        )
    normalized = normalize_answer(answer)
    exact_match = any(
        normalized == normalize_answer(reference)
        for reference in references
    )
    return RewardResult(
        float(exact_match),
        True,
        exact_match,
        answer,
    )

这段代码只评价最终任务结果,没有为“搜了几次”或“query 看起来不错”增加额外分数。对同一道问题的 $G$ 条完整轨迹,仍然使用 GRPO 组内基线:

$$A_i=R(\tau_i)-\frac{1}{G}\sum_{j=1}^{G}R(\tau_j).$$

Search-R1 论文同时讨论了 PPO 与 GRPO 等优化方式。本章选择 GRPO 的同题分组方案,让 Search-R1 与 8.1 节使用同一种 advantage 构造方法。

rollout_batch() 必须等整组轨迹全部结束后,才能执行下面的函数:

def assign_group_advantages(
    trajectories: list[Trajectory],
) -> int:
    groups: dict[int, list[Trajectory]] = {}
    for trajectory in trajectories:
        groups.setdefault(
            trajectory.question_index, []
        ).append(trajectory)

    degenerate = 0
    for group in groups.values():
        mean_reward = sum(
            item.reward for item in group
        ) / len(group)
        for item in group:
            item.advantage = (
                item.reward - mean_reward
            )
        if all(
            item.advantage == 0.0 for item in group
        ):
            degenerate += 1
    return degenerate

整组 reward 完全相同时,所有 advantage 都为 0,这一组不会产生有效梯度。训练代码会跳过这些轨迹,并记录退化组比例。组内均值不能等拆成 micro-batch 后再计算,否则比较基线已经被改变。

这个 advantage 会分配给轨迹中每一轮 Assistant 生成的 token。搜索 observation 只构成下一步决策所需的状态,不是策略做出的动作,因此 observation token 的 logprob 和 advantage 都设为 0:

full_tokens.extend(delta_observation)
old_logprobs_by_token.extend(
    [0.0] * len(delta_observation)
)
advantages_by_token.extend(
    [0.0] * len(delta_observation)
)

full_tokens.extend(turn.completion_tokens)
old_logprobs_by_token.extend(turn.logprobs)
advantages_by_token.extend(
    [trajectory.advantage]
    * len(turn.completion_tokens)
)

这种 observation mask 有两个作用。首先,训练器不会要求模型预测搜索引擎返回的文本;其次,observation 仍然保留在上下文中,后续 Assistant token 可以对这些证据进行条件建模。

搜索调用和最终回答共享一个轨迹级 advantage。若最终答案正确,导致成功的查询规划、搜索词和答案 token 都会得到鼓励;若最终答案错误,整条动作链都会被抑制。这正是 outcome reward 能够训练工具使用策略的原因。

8.3.6 构造多轮 PyTRIO Datum

train.pybuild_datum() 把前面几节的 token 连续性和 observation mask 汇合到一起。下面是完整的核心实现:

def build_datum(
    trajectory: Trajectory,
) -> TrainingDatum:
    if not trajectory.turns:
        raise ValueError(
            "不能用没有 assistant turn 的轨迹构造训练 Datum"
        )

    full_tokens: list[int] = []
    old_logprobs_by_token: list[float] = []
    advantages_by_token: list[float] = []
    assistant_token_count = 0

    for turn_index, turn in enumerate(
        trajectory.turns
    ):
        if len(turn.completion_tokens) != len(
            turn.logprobs
        ):
            raise ValueError(
                f"第 {turn_index + 1} 个 assistant turn "
                "的 token 与 logprob 长度不一致"
            )

        if turn_index == 0:
            delta_observation = turn.prompt_tokens
        elif turn.prompt_tokens[
            :len(full_tokens)
        ] == full_tokens:
            delta_observation = turn.prompt_tokens[
                len(full_tokens):
            ]
        else:
            raise ValueError(
                f"第 {turn_index + 1} 个 assistant turn "
                "的 prompt 不是已有轨迹的前缀扩展,"
                "无法安全对齐采样 logprob"
            )

        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)
        )
        assistant_token_count += len(
            turn.completion_tokens
        )

    if assistant_token_count == 0:
        raise ValueError(
            "不能用没有 assistant token 的轨迹构造训练 Datum"
        )
    if not (
        len(full_tokens)
        == len(old_logprobs_by_token)
        == len(advantages_by_token)
    ):
        raise ValueError(
            "完整轨迹的 token、logprob 和 advantage 长度不一致"
        )

    input_tokens = full_tokens[:-1]
    target_tokens = full_tokens[1:]
    old_logprobs = old_logprobs_by_token[1:]
    advantages = advantages_by_token[1:]
    if not (
        len(input_tokens)
        == len(target_tokens)
        == len(old_logprobs)
        == len(advantages)
    ):
        raise ValueError(
            "Datum 的 input、target、logprobs 和 advantages 长度不一致"
        )
    if len(input_tokens) > MAX_TRAIN_CONTEXT_TOKENS:
        raise ValueError(
            f"Datum 超过 {MAX_TRAIN_CONTEXT_TOKENS} token"
        )

    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,
            ),
        },
    )
    return TrainingDatum(datum, len(input_tokens))

第一轮的 delta_observation 是 system prompt、工具定义和用户问题;后续轮次的 delta_observation 是上一轮结束符、搜索结果以及新一轮 Assistant 前缀。它们都保留真实 target token,但 old logprob 与 advantage 为 0。只有 sampler 真正生成的 Assistant token 携带 rollout old logprob 和轨迹 advantage。

完整 group 的 reward 必须全部计算完毕,才能得到正确的组内均值。不能在单条轨迹结束后立刻更新策略,因为同组其他分支的奖励还未知。

多轮轨迹长度差异很大,直接把所有数据放入一个 padded batch 会浪费大量 token。配套代码根据样本长度把 Datum 装入多个 micro-batch,并按每批轨迹数占完整 rollout 轨迹数的比例重新加权 loss。这样可以在保持全局样本平均梯度语义的同时控制单批 padding 大小。

8.3.7 从 rollout 到一次参数更新

train.py 的外层循环把前面的模块串成一个完整训练 step。关键顺序不能交换:

batch = take_batch(
    examples,
    step * args.questions_per_batch,
    args.questions_per_batch,
)

sampling_client = (
    training_client
    .save_weights_and_get_sampling_client()
)

trajectories = rollout_batch(
    sampling_client=sampling_client,
    tokenizer=tokenizer,
    search_client=search_client,
    examples=batch,
    config=rollout_config,
)

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

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

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

第一,save_weights_and_get_sampling_client() 在每个 step 开始时导出当前 LoRA 权重,因此 old logprob 来自本次更新前的当前策略。第二,rollout_batch() 内部已经完成整组 reward 和 advantage,build_training_datums() 会跳过 advantage 全为 0 的轨迹。第三,所有 micro-batch 只累计梯度,整个 logical batch 最后只调用一次 optim_step()

PyTRIO 会对一次 forward_backward() 内的样本取平均。如果第 $k$ 个 micro-batch 有 $n_k$ 条轨迹,而完整 rollout batch 有 $N$ 条轨迹,代码会先把该批 advantage 乘以 $n_k/N$:

$$\sum_k\frac{n_k}{N}\operatorname{mean}(\mathcal L_k)=\operatorname{mean}(\mathcal L_{\mathrm{global}}).$$

对应实现只重新包装 Datum,不修改 target token 和 old logprob:

micro_batch_weight = np.float32(
    len(micro_batch) / total_samples
)
weighted_datums = []
for item in micro_batch:
    loss_inputs = item.datum.loss_fn_inputs
    weighted_datums.append(
        trio.Datum(
            model_input=item.datum.model_input,
            loss_fn_inputs={
                "target_tokens": (
                    loss_inputs["target_tokens"]
                    .to_numpy()
                ),
                "logprobs": (
                    loss_inputs["logprobs"]
                    .to_numpy()
                ),
                "advantages": (
                    loss_inputs["advantages"]
                    .to_numpy()
                    * micro_batch_weight
                ),
            },
        )
    )

因此,动态拆批只改变执行时的 padding 和显存占用,不改变逻辑 batch 的样本平均梯度。

8.3.8 运行、评测与实验边界

Search-R1 代码按数据、协议、环境、rollout 和训练职责拆分:

文件作用
prepare_data.py准备 NQ 与 HotpotQA 数据
data.py加载统一 JSONL 样本
protocol.py定义搜索工具与多轮消息协议
search.py封装搜索后端并统一结果格式
rollout.py执行多轮搜索状态机
reward.py抽取短答案并计算规则奖励
train.py构造 observation mask、micro-batch 并更新策略
eval.py在同一环境中评测 Base Model 或 checkpoint
analyse.py汇总评测输出

评测没有另写一套简化的搜索逻辑,而是复用同一个 rollout_batch()。主要区别是每道题只保留一条轨迹,不再需要训练阶段的组内比较:

config = RolloutConfig(
    group_size=1,
    max_search_calls=args.max_search_calls,
    max_assistant_turns=args.max_assistant_turns,
    max_trajectory_tokens=(
        args.max_trajectory_tokens
    ),
    max_assistant_tokens=args.max_assistant_tokens,
    max_tool_response_tokens=(
        args.max_tool_response_tokens
    ),
    search_concurrency=args.search_concurrency,
    temperature=args.temperature,
    top_p=args.top_p,
    seed=args.seed,
)

batch_trajectories = rollout_batch(
    sampling_client,
    tokenizer,
    search_client,
    batch,
    config,
)

model_path 留空时,eval.py 创建 Base Model sampler;传入 save_weights_for_sampler() 返回的路径时,则评测训练后的 checkpoint。两种模型仍使用相同的协议、搜索后端、预算、题集与答案判定器。

首先准备数据:

python docs/chapter8/search-r1/prepare_data.py

Wikipedia 后端无需 API Key,适合验证完整链路:

python docs/chapter8/search-r1/train.py \
    --max-steps 1 \
    --questions-per-batch 1 \
    --group-size 4 \
    --max-search-calls 2 \
    --search-backend wikipedia \
    --search-concurrency 3 \
    --swanlab-mode disabled

使用同一个 Wikipedia 环境评测 Base Model:

python docs/chapter8/search-r1/eval.py \
    --search-backend wikipedia \
    --batch-size 8 \
    --output docs/chapter8/search-r1/eval_result/base.jsonl

评测训练后的 sampler weights 时,保持其余参数不变,只增加 checkpoint 路径:

python docs/chapter8/search-r1/eval.py \
    --search-backend wikipedia \
    --batch-size 8 \
    --model-path 'trio://runxxxxxxxxxx' \
    --output docs/chapter8/search-r1/eval_result/checkpoint.jsonl

确认流程后,再逐步扩大问题数、group size、搜索次数与轨迹预算。实验至少应记录:

  • 最终答案 exact match;
  • 有效答案格式比例;
  • 平均搜索调用次数与零搜索比例;
  • 搜索失败率、超时率和 observation 截断率;
  • 平均轨迹长度与每步有效训练 token;
  • 退化组比例;
  • rollout、搜索和训练各阶段耗时。

搜索结果具有外部依赖。在线页面会更新,搜索服务的排序也可能变化。比较 Base Model 与训练 checkpoint 时,应固定题集、搜索后端、搜索预算、sampling 参数和评测时间窗口。若需要严格可复现的论文实验,可以将检索语料与索引固定为只读快照。

本章只更新语言模型策略,搜索后端保持固定。若检索器本身也参与学习,就需要额外定义检索动作空间、索引版本和跨模块 credit assignment,这已经超出本节示例范围。

Search-R1 展示了如何把 GRPO 的组内相对优势扩展到包含外部 observation 的多轮搜索轨迹。下一节的 ReTool 会保留同样的多轮 rollout 与 observation mask,将环境从搜索引擎换成代码解释器,并使用 PyTRIO 内置的 PPO 更新策略。

8.4 ReTool

代码来源: 本节依据作者开源仓库 agentic-rl-lab/05-retool 中的完整多文件实现整理。正文按照 prepare_data.pyprotocol.pysandbox.pyrollout.pyreward.pytrain.pyeval.pyanalysis.py 的真实调用链展开。

搜索工具帮助模型获取外部知识,代码解释器则帮助模型完成精确计算、符号操作与枚举验证。模型已经能在提示词或监督数据的引导下调用 Python 工具,但“会调用”还不等于“会在合适的时机调用”。有些问题口算更快,有些问题适合编写短程序,还有些代码执行失败后需要根据错误信息修正。

ReTool 使用强化学习训练这种策略性工具使用能力。模型自主决定是否调用代码解释器、编写什么代码、如何读取执行结果,以及何时给出最终答案。环境只根据最终数学答案提供结果奖励,成功轨迹中的工具决策会随整条轨迹共同得到强化。

安全提示: 配套 sandbox.py 只是带资源限制的本地 subprocess,不是可信安全沙箱。模型生成的代码仍能读取当前账号可访问的文件、继承环境变量并访问网络。真实训练应放在一次性容器、低权限虚拟机或专用沙箱服务中,并移除 PyTRIO、SwanLab、SSH 和云服务凭证。

ReTool 的代码链路与 Search-R1 结构相似,但环境从搜索服务换成了 Python 进程:

prepare_data.py / data.py
        ↓ 读取数学题与参考答案
protocol.py
        ↓ 解析 code_interpreter 调用
sandbox.py
        ↓ 执行代码,返回 stdout / stderr / timeout
rollout.py
        ↓ 追加执行结果并继续生成
reward.py
        ↓ 检查最后一个 \boxed{}
train.py
        ↓ feedback mask、PPO、optim_step
eval.py / analysis.py
        ↓ text-only 与 ReTool 对照评测

下面仍然顺着实际函数调用解释,而不是只看工具调用的文本格式。

8.4.1 从工具调用格式到工具使用策略

一次典型 ReTool 轨迹如下:

User: 计算满足条件的整数解数量。

Assistant -> code_interpreter(
    "count = ...\nprint(count)"
)
Tool: 37

Assistant: 根据枚举结果,答案为 \boxed{37}

如果代码出现错误,模型还可以利用 traceback 自我修正:

Assistant -> code_interpreter("print(sum(values))")
Tool: NameError: name 'values' is not defined

Assistant -> code_interpreter(
    "values = [...]\nprint(sum(values))"
)
Tool: 128

Assistant: \boxed{128}

从强化学习角度看,代码、执行错误、数值结果和最终回答共同构成轨迹。代码解释器是环境,Assistant 生成的工具调用与自然语言答案都是动作,stdout 或 stderr 是 observation。

ReTool 的多轮代码执行与结果奖励

图 8.5 ReTool 的多轮代码执行与结果奖励

ReTool 不需要为“调用了工具”单独加分。最终答案 reward 会通过轨迹级 advantage 传递到前面的代码动作。这样,只有真正帮助任务成功的工具使用方式会得到正向信号,冗余或误导性的调用会随失败轨迹受到抑制。

8.4.2 Cold-start、训练数据与工具协议

原始 ReTool 包含两个阶段:

  1. Cold-start SFT:使用带代码调用与执行反馈的轨迹,让模型先掌握协议和基本交互形式;
  2. Reinforcement Learning:模型自主 rollout,通过最终答案 reward 学习何时以及如何使用工具。

Cold-start 的作用是让初始策略产生一定比例的有效工具调用。如果基础模型完全不会输出工具协议,早期轨迹几乎全部失败,GRPO 很难从组内比较中获得有意义的探索信号。

本章代码选择已经具备原生 function calling 能力的 Qwen3.5,并直接展示第二阶段的 PyTRIO 训练,因此没有附带一条专门的 Cold-start SFT 流程。这是教学实现与原论文设置之间的重要边界。更换为缺少工具调用能力的基础模型时,应先准备 SFT 数据,验证合法调用率达到可训练水平,再启动 RL。

RL 数据来自 DAPO-Math-17k。原始题目外面带有要求 Answer: $Answer 的统一模板,它会与本节要求的 \boxed{} 格式冲突,因此 prepare_data.py 先剥掉外层模板,再只保留题目和参考答案:

def normalize_row(
    index: int,
    row: dict[str, Any],
) -> dict[str, Any] | None:
    prompt = row.get("prompt")
    if not isinstance(prompt, list) or not prompt:
        return None

    question = strip_dapo_template(
        str(prompt[0].get("content") or "")
    )
    reward_model = row.get("reward_model") or {}
    ground_truth = (
        reward_model.get("ground_truth")
        if isinstance(reward_model, dict)
        else None
    )
    if isinstance(ground_truth, list):
        ground_truth = (
            ground_truth[0] if ground_truth else None
        )
    answer = str(ground_truth or "").strip()
    if not question or not answer:
        return None

    return {
        "id": str(
            row.get("extra_info", {}).get("index")
            or index
        ),
        "question": question,
        "answer": answer,
        "data_source": str(
            row.get("data_source") or "dapo_math"
        ),
    }

和 Search-R1 一样,训练样本不含标准代码和工具调用次数。策略必须自己探索是否调用解释器以及生成什么代码。

代码解释器使用结构化工具定义:

CODE_TOOL = {
    "type": "function",
    "function": {
        "name": "code_interpreter",
        "description": (
            "Run Python code and return printed output. "
            "Each execution starts fresh."
        ),
        "parameters": {
            "type": "object",
            "properties": {
                "code": {"type": "string"},
            },
            "required": ["code"],
        },
    },
}

protocol.py 对工具调用同样采用严格解析。Python 代码可以包含比较符、引号和多行文本,因此这里只校验调用数量、调用位置以及代码是否为空:

def parse_assistant(text: str) -> ParsedAssistant:
    matches = list(TOOL_CALL_PATTERN.finditer(text))
    if not matches:
        kind = (
            "invalid" if "<tool_call>" in text else "answer"
        )
        return ParsedAssistant(
            kind=kind,
            content=text.strip(),
        )

    if len(matches) != 1 or text[matches[0].end():].strip():
        return ParsedAssistant(
            kind="invalid",
            content=text.strip(),
        )

    code = matches[0].group(1).strip()
    if not code:
        return ParsedAssistant(
            kind="invalid",
            content=text.strip(),
        )

    content = text[:matches[0].start()].strip()
    return ParsedAssistant(
        kind="tool",
        content=content,
        code=code,
    )

合法输出只允许三种去向:tool 进入代码执行环境,answer 结束轨迹并接受结果奖励,invalid 结束轨迹并按错误答案处理。系统不为格式接近但不完整的工具调用做猜测性修复,否则训练环境中的动作定义会变得不稳定。

rollout 状态机与 Search-R1 相似,只是环境动作从 search(query) 换成 code_interpreter(code)。每条轨迹都受到以下预算约束:

  • 最大代码调用次数;
  • 最大 Assistant 轮数;
  • 单轮生成与执行结果 token 数;
  • 完整轨迹 token 数;
  • 单次执行时间;
  • 本地执行器并发数。

同一次代码调用不保留 Python 变量和进程状态。模型需要在每次调用中重新定义所需数据,这使轨迹语义更清楚,也减少了不同执行之间的隐式依赖。

8.4.3 执行结果怎样接回真实 token 轨迹

ReTool 也不能把 Assistant 文本 decode 后重新 tokenize。protocol.py 使用一个占位 Assistant 提取模板闭合 token,再把实际执行结果对应的 observation token 接到真实 completion 后:

def build_next_prompt(
    tokenizer: Any,
    messages_before_assistant: list[dict[str, Any]],
    previous_prompt_tokens: list[int],
    completion_tokens: list[int],
    next_tool_message: dict[str, Any],
) -> list[int]:
    canonical_prompt = build_prompt(
        tokenizer,
        messages_before_assistant,
    )
    placeholder_message = {
        "role": "assistant",
        "content": "x",
    }
    messages_with_assistant = [
        *messages_before_assistant,
        placeholder_message,
    ]
    canonical_assistant_end = _render_chat(
        tokenizer,
        messages_with_assistant,
        add_generation_prompt=False,
    )

    placeholder_tokens = _encoded_text_tokens(
        tokenizer, "x"
    )
    canonical_action = [
        *canonical_prompt,
        *placeholder_tokens,
    ]
    if canonical_assistant_end[
        :len(canonical_action)
    ] != canonical_action:
        raise ValueError(
            "chat template 无法定位 assistant 结束边界"
            "(占位内容也不匹配)"
        )
    assistant_closing_tokens = canonical_assistant_end[
        len(canonical_action):
    ]

    canonical_next_prompt = build_prompt(
        tokenizer,
        [*messages_with_assistant, next_tool_message],
    )
    if canonical_next_prompt[
        :len(canonical_assistant_end)
    ] != canonical_assistant_end:
        raise ValueError(
            "加入 tool observation 后 chat template 改写了历史消息"
        )
    observation_tokens = canonical_next_prompt[
        len(canonical_assistant_end):
    ]

    overlap = _suffix_prefix_overlap(
        completion_tokens,
        assistant_closing_tokens,
    )
    return [
        *previous_prompt_tokens,
        *completion_tokens,
        *assistant_closing_tokens[overlap:],
        *observation_tokens,
    ]

这里使用占位内容,是因为 Qwen chat template 可能把真实采样文本中的 reasoning 与 content 重新分段。占位消息只用于计算固定模板边界,真实模型动作仍始终来自 completion_tokens。下一轮 prompt_tokens 必须以前一轮的完整 token 为前缀,否则 train.py 会拒绝构造 Datum

8.4.4 Outcome reward 与 interpreter feedback mask

本章从最终回答的末尾抽取最后一个 $\boxed{}$,使用 math_verify 判断预测值与参考答案是否数学等价:

def extract_last_boxed(text: str) -> str | None:
    marker = "\\boxed{"
    index = text.rfind(marker)
    if index < 0:
        return None

    start = index + len(marker)
    depth = 1
    for position in range(start, len(text)):
        if text[position] == "{":
            depth += 1
        elif text[position] == "}":
            depth -= 1
            if depth == 0:
                return text[start:position]
    return None


def answers_equivalent(
    prediction: str,
    reference: str,
) -> bool:
    try:
        return bool(
            verify(
                parse(f"${reference}$"),
                parse(f"${prediction}$"),
            )
        )
    except Exception:
        return False


def score_answer(
    text: str,
    reference: str,
) -> RewardResult:
    answer = extract_last_boxed(
        text[-ANSWER_WINDOW_CHARS:]
    )
    if answer is None:
        return RewardResult(
            -1.0, False, False, None
        )

    correct = answers_equivalent(
        answer.strip(),
        reference.strip(),
    )
    return RewardResult(
        1.0 if correct else -1.0,
        correct,
        True,
        answer,
    )

奖励函数为:

$$R(\tau)=\begin{cases}+1, & \text{最终答案正确},\ -1, & \text{答案错误或格式无效}.\end{cases}$$

再对同题轨迹计算中心化 advantage:

$$A_i=R(\tau_i)-\bar R.$$

代码执行结果属于环境 observation。它应出现在模型上下文中,同时不能进入策略 loss。ReTool 论文将这一处理称为 interpreter feedback mask,本章采用与 Search-R1 相同的 token 对齐方式:

token 区域是否保留在上下文old logprobadvantage
初始 prompt00
Assistant 推理与代码调用rollout logprob轨迹 advantage
stdout、stderr 或超时信息00
最终答案rollout logprob轨迹 advantage

配套代码使用 PyTRIO 的 PPO loss:

PPO_LOSS_CONFIG = {
    "clip_low_threshold": 0.8,
    "clip_high_threshold": 1.28,
}

training_client.forward_backward(
    datums,
    loss_fn="ppo",
    loss_fn_config=PPO_LOSS_CONFIG,
).result()

原始 ReTool 论文的强化学习阶段使用 PPO。本章保留多轮代码轨迹、结果奖励和 interpreter feedback mask,同时加入同题分组采样与中心化 advantage;因此,训练链路可以概括为“GRPO 构造相对优势,PPO clipping 控制策略更新”。

这里采用非对称概率比边界。旧策略概率比低于 0.8 或高于 1.28 时,PPO 会限制对应更新。无论使用何种 clip 参数,rollout sampler 都需要在每个训练 step 开始前从当前权重刷新。

8.4.5 代码执行环境与安全边界

执行模型生成的代码具有真实风险。本章 sandbox.py 使用本地 subprocess,并加入了以下工程限制:

  • 使用参数列表启动 Python,不经过 shell;
  • 每次调用创建独立进程;
  • 设置 wall-clock timeout 与 CPU 时间上限;
  • 超时后终止整个子进程组;
  • 限制 BLAS 等数值库的线程数;
  • stdout 和 stderr 写入临时文件,并限制读回长度;
  • 限制并发执行数量;
  • 清理可能破坏 chat template 的特殊标记。

执行入口是 LocalPythonSandbox.run_code()。模型代码先被嵌入一个设置 CPU 限制的 bootstrap,再通过参数列表传给新的 Python 进程:

def run_code(self, code: str) -> ExecResult:
    started = time.perf_counter()
    self.stats.calls += 1
    bootstrap = _BOOTSTRAP.format(
        cpu=int(self.timeout) + 5,
        code=repr(code),
    )
    env = {
        **os.environ,
        **_CHILD_ENV_OVERRIDES,
    }

    with self._semaphore:
        with (
            tempfile.TemporaryFile() as stdout_file,
            tempfile.TemporaryFile() as stderr_file,
        ):
            process = subprocess.Popen(
                [
                    self.python,
                    "-B",
                    "-c",
                    bootstrap,
                ],
                stdout=stdout_file,
                stderr=stderr_file,
                env=env,
                start_new_session=True,
            )
            timed_out = False
            try:
                returncode = process.wait(
                    timeout=self.timeout
                )
            except subprocess.TimeoutExpired:
                timed_out = True
                os.killpg(
                    process.pid,
                    signal.SIGKILL,
                )
                returncode = None
                process.wait()

            stdout = _truncate_tail(
                _read_capped(stdout_file)
            )
            stderr = _truncate_tail(
                _read_capped(stderr_file)
            )

    latency = time.perf_counter() - started
    ok = not timed_out and returncode == 0
    if timed_out:
        self.stats.timeouts += 1
    elif ok:
        self.stats.successes += 1
    else:
        self.stats.errors += 1
    self.stats.latency_total += latency
    return ExecResult(
        ok,
        stdout,
        stderr,
        returncode,
        timed_out,
        latency,
    )

stdout 和 stderr 写入临时文件而不是 PIPE,父进程只读取文件尾部的有限字节,因此模型即使持续打印,也不会把全部输出保存在父进程内存中。超时时使用 start_new_session=True 创建的独立进程组,随后由 killpg() 一次终止整组子进程。

执行结果被统一转换为一条 tool observation:

def format_tool_content(
    self,
    result: ExecResult,
) -> str:
    if result.timed_out:
        return (
            "Error: execution timed out after "
            f"{self.timeout:.0f} seconds."
        )
    if result.ok:
        return sanitize_tool_content(result.stdout)
    if result.stderr:
        return sanitize_tool_content(result.stderr)
    return (
        "Error: process exited with code "
        f"{result.returncode}."
    )

成功时返回 stdout,语法错误和运行时异常返回原始 stderr,超时返回固定错误文本。当前 run_code() 不会自动替模型补 print(),所以 system prompt 明确要求模型打印希望观察的值;没有打印内容时,成功 observation 就是空字符串。

这些措施主要用于控制意外的死循环、输出洪水和资源消耗。它们不能构成可信安全隔离。本地子进程仍可能读取运行账号可访问的文件、访问网络、读取继承的环境变量,或利用宿主环境中的漏洞。

因此,运行 ReTool 前必须遵守以下安全要求:

  1. 不要在含有生产凭证、私钥或敏感数据的主机上直接执行模型代码;
  2. 使用一次性容器、低权限虚拟机或专用沙箱服务;
  3. 默认关闭网络,并设置只读文件系统与最小化可见目录;
  4. 移除无关环境变量和云服务凭证;
  5. 为 CPU、内存、进程数、磁盘和执行时间设置硬限制;
  6. 保存代码、退出状态与资源指标,便于审计异常轨迹。

工具描述中的“independent execution”表示每次调用不共享 Python 状态,它不代表宿主机级安全隔离。教学试跑也应先放入隔离环境。

8.4.6 多轮代码 rollout 怎样推进

rollout.py 分开保存可训练的 Assistant 动作、完整轨迹状态和“已经生成、尚未执行”的代码调用:

@dataclass
class AssistantTurn:
    prompt_tokens: list[int]
    completion_tokens: list[int]
    logprobs: list[float]
    text: str


@dataclass
class Trajectory:
    example: MathExample
    group_index: int
    messages: list[dict[str, Any]]
    next_prompt_tokens: list[int] | None = None
    question_index: int = 0
    turns: list[AssistantTurn] = field(
        default_factory=list
    )
    code_calls: int = 0
    final_text: str = ""
    reward: float = -1.0
    advantage: float = 0.0
    valid_format: bool = False
    correct: bool = False
    done: bool = False


@dataclass(frozen=True)
class PendingExecution:
    trajectory: Trajectory
    code: str
    call_id: str
    messages_before_assistant: list[dict[str, Any]]
    assistant_text: str
    prompt_tokens: list[int]
    completion_tokens: list[int]

begin_advance() 消费一次 Assistant 输出,并决定轨迹是结束还是进入执行环境:

def begin_advance(
    trajectory: Trajectory,
    prompt_tokens: list[int],
    sequence: Any,
    tokenizer: Any,
    config: RolloutConfig,
) -> PendingExecution | None:
    tokens, logprobs, text = read_sequence(
        sequence, tokenizer
    )
    text = text.strip()
    trajectory.turns.append(
        AssistantTurn(
            prompt_tokens,
            tokens,
            logprobs,
            text,
        )
    )
    parsed = parse_assistant(text)

    can_code = (
        parsed.kind == "tool"
        and trajectory.code_calls
        < config.max_code_calls
        and len(trajectory.turns)
        < config.max_assistant_turns
    )
    if not can_code:
        trajectory.messages.append(
            {"role": "assistant", "content": text}
        )
        trajectory.final_text = text
        trajectory.done = True
        return None

    call_id = (
        f"code-{trajectory.question_index}-"
        f"{trajectory.group_index}-"
        f"{trajectory.code_calls + 1}"
    )
    messages_before_assistant = list(
        trajectory.messages
    )
    trajectory.messages.append(
        {"role": "assistant", "content": text}
    )
    return PendingExecution(
        trajectory,
        parsed.code or "",
        call_id,
        messages_before_assistant,
        text,
        prompt_tokens,
        tokens,
    )

同一 round 内可能有许多轨迹同时请求执行代码。它们可以并发,但每条轨迹仍必须等待自己的 observation 返回以后才能开始下一轮生成:

async def execute_pending_async(
    pendings: list[PendingExecution],
    sandbox: LocalPythonSandbox,
) -> list[ExecResult]:
    return list(
        await asyncio.gather(
            *(
                sandbox.arun_code(p.code)
                for p in pendings
            )
        )
    )


async def advance_round(
    responses: list[Any],
    requests: list[SampleRequest],
    trajectories: list[Trajectory],
    tokenizer: Any,
    sandbox: LocalPythonSandbox,
    config: RolloutConfig,
    progress_callback: Callable[[int], None] | None,
) -> None:
    pendings: list[PendingExecution] = []
    for request, response in zip(
        requests, responses, strict=True
    ):
        trajectory = trajectories[
            request.trajectory_index
        ]
        pending = begin_advance(
            trajectory,
            request.prompt_tokens,
            response.sequences[0],
            tokenizer,
            config,
        )
        if pending is not None:
            pendings.append(pending)
        elif (
            trajectory.done
            and progress_callback is not None
        ):
            progress_callback(1)

    if not pendings:
        return
    results = await execute_pending_async(
        pendings, sandbox
    )
    for pending, result in zip(
        pendings, results, strict=True
    ):
        finish_advance(
            pending,
            result,
            tokenizer,
            sandbox,
            config,
        )
        if (
            pending.trajectory.done
            and progress_callback is not None
        ):
            progress_callback(1)

finish_advance() 调用上一节的 format_tool_content(),再用 8.4.3 节的 build_next_prompt() 构造连续 token。如果执行结果过长,fit_tool_content() 会反复保留尾部,因为 traceback 和最终数值通常位于末尾;结果仍放不进轨迹预算时,本条轨迹直接结束。

首轮分支方式与 Search-R1 相同:每个问题只提交一个 prompt,并设置 num_samples=group_size。每条返回序列通过 copy.deepcopy(root) 变成独立轨迹。执行结果返回以后,各分支的 prompt 已经不同,后续请求全部设置 num_samples=1,直到输出最终答案或耗尽代码调用、轮次与 token 预算。

所有轨迹结束以后,状态机才依次调用 score_trajectory()assign_group_advantages()。因此,一条轨迹的 traceback 可以影响它后续生成的代码,却不会越过轨迹边界影响同组其他分支。

8.4.7 使用 PyTRIO 训练 ReTool

首先,使用当前权重创建采样客户端,并生成同题多轮轨迹:

sampling_client = (
    training_client.save_weights_and_get_sampling_client()
)

trajectories = rollout_batch(
    sampling_client=sampling_client,
    tokenizer=tokenizer,
    sandbox=sandbox,
    examples=batch,
    config=rollout_config,
)

rollout_batch() 完成以下工作:

  1. 第一轮对同一问题采样 group_size 个分支;
  2. 解析每条分支中的 code_interpreter 调用;
  3. 并发执行模型代码并追加 tool observation;
  4. 对仍未结束的轨迹继续采样;
  5. 抽取最终答案,计算 reward 与组内 advantage。

每个训练 step 都重新调用 save_weights_and_get_sampling_client(),保证 rollout old logprob 来自更新前的当前策略。如果同题所有轨迹都答对或都答错,中心化 advantage 全为 0,build_training_datums() 会跳过这一组;没有剩余 Datum 时,本 step 不执行 optim_step()

随后,将每条完整轨迹构造为一个 Datum。ReTool 的 build_datum() 与 8.3.6 节遵循同一套前缀扩展规则,核心循环如下:

for turn_index, turn in enumerate(trajectory.turns):
    if turn_index == 0:
        delta_observation = turn.prompt_tokens
    elif turn.prompt_tokens[
        :len(full_tokens)
    ] == full_tokens:
        delta_observation = turn.prompt_tokens[
            len(full_tokens):
        ]
    else:
        raise ValueError(
            "下一轮 prompt 不是已有轨迹的前缀扩展"
        )

    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
        ),
    },
)

delta_observation 中包含初始题目、Assistant 结束符以及 stdout、stderr 或 timeout 文本,它们的 advantage 全为 0;Assistant 推理、代码调用与最终答案共享轨迹 advantage。这就是 interpreter feedback mask 的实际代码形态。

由于轨迹长度差异较大,代码再进行动态 micro-batch 装箱:

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

for micro_batch in micro_batches:
    weighted = weight_micro_batch_for_global_mean(
        micro_batch,
        total_samples=len(trajectories),
    )
    training_client.forward_backward(
        weighted,
        loss_fn="ppo",
        loss_fn_config=PPO_LOSS_CONFIG,
    ).result()

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

PyTRIO 会对单次 forward_backward() 中的样本取平均。不同 micro-batch 大小不一致时,直接累积会放大小批次的权重。本章按照每个 micro-batch 样本数占全局轨迹数的比例缩放 advantage:

$$\sum_k \frac{n_k}{N}\mathrm{mean}(\mathcal L_k) = \mathrm{mean}(\mathcal L_{\mathrm{global}}).$$

这样,动态拆批只改变执行方式,不改变全局样本平均的梯度含义。

8.4.8 运行与评测 ReTool

配套目录包含以下文件:

文件作用
prepare_data.py准备 DAPO-Math-17k 训练集与 dev 集
data.py加载统一数学样本
protocol.py定义代码工具和多轮消息协议
sandbox.py执行 Python 代码并记录资源指标
rollout.py执行多轮“生成—运行—观察”状态机
reward.py抽取 $\boxed{}$ 并做数学等价判定
train.py构造 feedback mask、micro-batch 并执行 PPO
eval.py统一评测 text-only 与 ReTool 模式
analysis.py汇总不同 checkpoint 的指标

prepare_data.py 生成的 DAPO-Math-17k train.jsonl 用于 RL,并另外切出 50 道 dev.jsonl 做快速检查。正式的 eval.py 使用固定 30 道 AIME 2025,避免直接在训练题上报告结果。

在 ReTool 模式下,每道题的 val_n 条候选仍由同一个多轮状态机生成:

config = RolloutConfig(
    group_size=args.val_n,
    max_code_calls=args.max_code_calls,
    max_assistant_turns=args.max_assistant_turns,
    max_trajectory_tokens=(
        args.max_trajectory_tokens
    ),
    max_assistant_tokens=args.max_assistant_tokens,
    max_tool_response_tokens=(
        args.max_tool_response_tokens
    ),
    temperature=args.temperature,
    top_p=args.top_p,
    seed=args.seed,
)

trajectories = await rollout_batch_async(
    sampling_client,
    tokenizer,
    sandbox,
    examples,
    config,
)

text-only 模式则不注册代码工具,直接对同一题采样 val_n 个单轮答案。两种模式最终都写出逐题 JSONL,并汇总:

  • Average@N:全部生成结果中的平均正确率;
  • Pass@N:至少有一条候选答对的问题比例;
  • format_rate:能够抽取合法 \boxed{} 的比例;
  • mean_code_callsmean_turns:工具成本和轨迹长度。

因此,比较 checkpoint 时至少需要固定题集、val_n、temperature、top-p、最大轨迹长度和答案抽取逻辑;比较 text-only 与 ReTool 时,还要额外报告执行延迟和代码调用数。

首先准备数据:

python docs/chapter8/retool/prepare_data.py

确认代码执行器位于隔离环境后,再进行一步试跑:

python docs/chapter8/retool/train.py \
    --max-steps 1 \
    --questions-per-batch 1 \
    --group-size 4 \
    --max-code-calls 2 \
    --sandbox-workers 2 \
    --swanlab-mode disabled

准备好 AIME 2025 本地数据后,可以先跑无工具对照,再跑相同采样参数的 ReTool 评测:

python docs/chapter8/retool/eval.py \
    --mode text-only \
    --val-n 12 \
    --temperature 1.0 \
    --top-p 0.7 \
    --output docs/chapter8/retool/eval-results/aime25-text-only-base.jsonl

python docs/chapter8/retool/eval.py \
    --mode retool \
    --val-n 12 \
    --temperature 1.0 \
    --top-p 0.7 \
    --output docs/chapter8/retool/eval-results/aime25-retool-base.jsonl

评测训练后的模型时,为两条命令增加相同的 --model-path trio://...。只有 Base Model、checkpoint、text-only 和 ReTool 使用同一套采样设置时,差异才可以归因于训练或工具环境。

训练指标应同时覆盖任务、策略和环境三个层面:

  • 任务结果:正确率、合法 $\boxed{}$ 比例、退化组比例;
  • 工具行为:代码调用率、每条轨迹平均调用次数、零调用正确率;
  • 自我修正:首次执行报错后重试并最终答对的比例;
  • 环境状态:执行成功率、异常率、超时率、平均 latency;
  • 策略稳定性:PPO probability ratio、clip fraction、response length;
  • 系统效率:rollout、执行器、远端训练和 checkpoint 各阶段耗时。

评测时可以分别运行 text-only 与 ReTool 模式。两者需要使用同一批问题、相同采样参数和相同答案判定器,才能看出工具环境带来的净收益。工具模式还应报告额外延迟与执行成本,因为准确率提升通常伴随更多环境调用。

至此,我们已经从 GRPO 的单轮结果奖励、OPD 的逐 token Teacher 反馈,走到了 Search-R1 与 ReTool 的多轮环境交互。下面的本章小结会把四个模块重新收束到同一条“rollout—反馈—advantage—策略更新”数据流中。

本章小结

本章从 GRPO 出发,逐步建立了一条可以扩展到 Agent 的大模型强化学习路径。

GRPO 使用同题多次采样构造组内相对优势,省去了单独训练 Value Model 的过程。它把算法拆成清晰的三层:策略生成 rollout,环境给出 reward,更新 loss 根据新旧策略概率比优化模型。正确的 reward、token 对齐和 on-policy sampler 刷新共同决定训练是否有效。

OPD 保留 on-policy rollout,将规则奖励替换为 Teacher 对 Student 轨迹的逐 token 概率反馈。reverse KL 为每个动作提供稠密信号,使小模型可以在自身会访问的状态上学习 Teacher 能力。Teacher 与 Student 的 tokenizer、推理模式和能力互补关系需要在训练前验证。

Search-R1 与 ReTool 进一步把轨迹扩展为多轮环境交互。搜索结果和代码执行结果进入上下文,却通过 observation mask 排除在策略 loss 之外;模型生成的搜索词、代码、推理和最终答案共享轨迹级结果信号。两者都复用了 GRPO 的分组采样与相对优势,只替换了工具协议、环境实现和 reward。

四个示例最终汇合为同一条 PyTRIO 数据流:

$$\text{当前策略}\rightarrow\text{多样化 rollout}\rightarrow\text{环境或 Teacher 反馈}\rightarrow\text{token-level advantage}\rightarrow\text{策略更新}\rightarrow\text{新策略}.$$

掌握这条数据流后,新的 Agentic RL 任务可以从五个问题开始设计:

  1. Agent 可以执行哪些动作,环境会返回什么 observation?
  2. 一条轨迹在什么条件下结束,预算如何限制?
  3. reward 能否稳定反映最终任务目标?
  4. 哪些 token 属于策略动作,哪些 token 必须被 mask?
  5. rollout、旧 logprob、advantage 与当前权重是否严格对齐?

当这五个问题都有可验证的答案时,一个工具 Demo 才具备进入强化学习训练的基础。

继续学习: 受篇幅和章节定位限制,本章集中介绍了 GRPO、OPD、Search-R1 与 ReTool,仍有许多 Agentic RL 方法和交互环境无法逐一展开。如果你希望继续学习 OPSD、DAPO、GSPO、ALFWorld 等内容,可以前往作者持续维护的 agentic-rl-lab。仓库将继续整理相关算法的原理、可运行代码与真实训练过程。

参考资料

  1. Shao, Z. et al. DeepSeekMath: Pushing the Limits of Mathematical Reasoning in Open Language Models. 2024.
  2. Schulman, J. et al. Proximal Policy Optimization Algorithms. 2017.
  3. Agarwal, R. et al. On-Policy Distillation of Language Models: Learning from Self-Generated Mistakes. 2023.
  4. Li, Y. et al. Rethinking On-Policy Distillation of Large Language Models: Phenomenology, Mechanism, and Recipe. 2026.
  5. Sun, J. et al. EasyOPD: An Easy-to-use On-Policy Distillation Framework for Large Language Models. 2026.
  6. Jin, B. et al. Search-R1: Training LLMs to Reason and Leverage Search Engines with Reinforcement Learning. 2025.
  7. Feng, J. et al. ReTool: Reinforcement Learning for Strategic Tool Use in LLMs. 2025.
  8. PyTRIO. Loss Function Guide.
  9. PyTRIO. Search-R1 Example.
  10. Agentic-RL Lab (不要葱姜蒜). agentic-rl-lab.
《从零开始构建大模型》」· 第 10/10 章 · 内容开源自 datawhalechina/happy-llm CC BY-NC-SA),版权归原作者