PyTRIO快速入门实战篇(二):用 GRPO 提升 GSM8K 数学推理准确率

用 PyTRIO 对 Qwen3.5-4B 进行 50 step 的 GRPO 训练。在 GSM8K test split 的前 100 道题上,模型的准确率从 79% 提升到 95%,训练与评估总成本仅 1.7 元。

在上一篇实战中,我们用监督微调让模型学习标准答案。这一篇换一种思路:不给模型逐 token 的标准答案,而是让模型自己生成多种解法,再根据答案好坏获得奖励。

本节,我们将使用 PyTRIO 在 GSM8K 上完成一次 GRPO 训练,并通过同一套评估脚本比较 base model 与训练后模型的数学推理准确率。

话不多说,我们先看结果:

模型 GRPO step 评估样本 准确率
Qwen3.5-4B Base 0 100 79%
Qwen3.5-4B GRPO 50 100 95%

经过50个step的GRPO训练后,Qwen3.5-4B在79%的准确率提升到95%的准确率。

50step的训练消耗了0.29M的训练Token,评估消耗了0.06M Token,所用的成本仅为1.7元,可以说很有性价比了:

训练 Token 评估 Token 总 Token 训练与评估总成本
0.29M 0.06M 0.35M 1.7 元

任务介绍

GSM8K 是一个小学数学应用题数据集。每条数据包含一道自然语言问题,以及带有解题过程和最终答案的参考结果。

例如,题目可能会问:

Janet 的鸭子每天产 16 个蛋。她每天早餐吃 3 个,并用 4 个做松饼,剩下的每个卖 2 美元。她每天能赚多少钱?

一个符合要求的模型模拟回答可以是:

text 复制代码
Janet 每天一共产 16 个蛋。
吃掉 3 个、做松饼用掉 4 个后,还剩:
16 - 3 - 4 = 9 个。

每个蛋卖 2 美元,所以每天的收入是:
9 × 2 = 18 美元。

因此,Janet 每天能赚 \boxed{18} 美元。

ps:这段回答只是为了展示模型预期的输出结构,并不是从某次 rollout 日志中截取的真实结果。

可以看到,模型需要写出一段可读的计算过程,并把最终数值放进 \boxed{18};评估脚本不会逐句判断推理过程,而是提取回答中最后一个 \boxed{...}。如果回答没有 boxed 答案,脚本会退而提取最后一个数字,再与 GSM8K 标准答案进行数值比较。

GRPO 在做什么

GRPO 的全称是 Group Relative Policy Optimization,由 DeepSeek 在 2024 年发布的论文 《DeepSeekMath: Pushing the Limits of Mathematical Reasoning in Open Language Models》 中提出。

在此之前,大模型强化学习常使用 PPO。PPO 属于 actor-critic 方法:除了正在训练的策略模型,还需要训练一个 value model,也就是 critic,用它估计每个状态的价值并计算 advantage。原论文指出,这个 value model 通常与策略模型规模相当,会额外带来明显的显存和计算负担;而大模型的 reward 往往只在回答结束时给出,这也增加了为每个 token 学准 value function 的难度。

GRPO 的解决方式是不再训练额外的 critic model。它让当前策略针对同一道题生成一组回答,再根据这组回答的 reward 估计 baseline,并判断每个回答相对组内水平是更好还是更差。这样既保留了 advantage 带来的相对训练信号,又减少了 PPO 中 value model 所需的训练资源。

回到这次实战,它的核心流程可以概括为四步:

  1. 对同一道题采样一组答案。
  2. 用可验证的 reward 函数为每个答案打分。
  3. 用"当前答案 reward - 同组平均 reward"得到 advantage。
  4. 提高组内高分答案的概率,降低低分答案的概率。

这里的"相对"很重要。为了方便说明,我们先用 4 个回答举例:假设它们的得分为 [1.0, 1.0, 0.2, 0.0],组内平均分是 0.55,那么对应 advantage 就是 [0.45, 0.45, -0.35, -0.55]。模型不需要额外的 critic,而是直接从同组答案的比较中得到训练信号。本次实际训练的 group size 是 16,也就是每道题会同时采样 16 个回答进行组内比较。

如果一组答案的 reward 完全相同,它们的 advantage 都是 0,无法提供相对优劣信息。训练脚本会跳过这样的 group。

为什么使用 PyTRIO

PyTRIO 是 TRIO 远程大模型后训练和推理服务的 Python SDK。简单来说,我们在本地用 Python 编写数据处理、reward 和训练循环,而模型采样、前向反向传播、优化器更新及权重保存等计算密集型任务,都由 PyTRIO 的远程服务完成。

这种分工很适合想学习大模型后训练、但手边没有 GPU 集群的开发者。我们不需要先配置显卡环境、部署推理服务或搭建分布式训练系统,只需要安装 pytrio、运行 trio login 完成登录,就可以从普通的本地 Python 环境发起模型训练。本地电脑主要负责控制实验流程,因此这次 Qwen3.5-4B 的 GRPO 训练也不要求本地 GPU。

PyTRIO 不只是把训练放到远程,它还把训练和采样放进了同一套 SDK。在这篇文章的代码里,我们先通过 ServiceClient 创建 LoRA TrainingClient,每个 step 再从当前训练权重得到 SamplingClient;模型完成一组 rollout 后,本地代码计算 reward 和 advantage,并构造 Datum 交回远程 trainer 更新参数。训练结束后,我们可以直接保存 sampler 权重,并用同一模型路径启动评估。

这条链路尤其适合 GRPO。因为 GRPO 既需要反复采样多个回答,又需要根据回答结果立即更新模型,如果训练和推理分别使用两套基础设施,实验代码和环境管理都会更复杂。使用 PyTRIO 后,我们可以把注意力集中在真正影响效果的部分------prompt、reward、group-relative advantage 和训练参数------并用一份 Python 脚本完成从 rollout 到权重保存的闭环。

从本次实验的结果看,这套方式也足够轻量:50 step 训练加评估共使用 0.35M Token,实际成本为 1.7 元。对于第一次尝试 RL 后训练的读者,它提供了一条不需要先购买硬件、同时又能完整理解 GRPO 数据流的实践路径。

准备工作

由于PyTRIO不挑设备,所以不需准备带有GPU的机器,我是在我的Macbook上完成的。

先进入示例目录并安装依赖:

bash 复制代码
cd pytrio-quick-start
python -m pip install -U pytrio transformers datasets numpy addict

然后登录 PyTRIO:

bash 复制代码
trio login

核心文件结构如下:

text 复制代码
pytrio-quick-start/
├── train.py  # GSM8K 数据加载、GRPO rollout、reward 与训练
└── eval.py   # base model 与 GRPO checkpoint 的异步评估

数据不需要手动下载。脚本第一次运行时,会通过 Hugging Face datasets 自动加载 openai/gsm8kmain 配置。

在开始训练前,可以先跑出 base model 的基线:

bash 复制代码
python eval.py --limit 100

这条命令使用 GSM8K test split 的前 100 道题,默认设置 temperature=0.0max_tokens=512。本次实测得到的 base model exact accuracy 为 79%

开始训练

使用下面的命令启动 50 step GRPO 训练:

bash 复制代码
python train.py --steps 50 --batch-size 2 --group-size 16 --max-tokens 512 --eval-limit 0

--eval-limit 0 表示训练结束后暂时跳过脚本内置的小规模评估。稍后我们会用独立的 eval.py,在完整的 100 条样本口径上比较结果。

看到下面的打印时,代表训练已经跑起来了:

本次训练采用的主要配置如下:

配置 本次运行值 作用
base model Qwen/Qwen3.5-4B 初始策略模型
LoRA rank 16 训练低秩适配器
数据加载范围 train split 前 200 条 候选训练题目
实际训练题数 前 100 条 50 step × 每步 2 道题
steps 50 参数更新次数
batch size 2 每个 step 使用的题目数
group size 16 每道题采样的回答数
rollout temperature 1.0 保持采样多样性
max tokens 512 单条回答的最大生成长度
learning rate 4e-5 Adam 学习率
seed 42 rollout 采样种子

脚本默认加载 train split 的前 200 条数据,但本次设置为 50 step、每步 2 道题,因此实际依次使用其中的前 100 道题。每道题最多产生 16 条 rollout,也就是每个 step 最多产生 32 条、本次训练最多产生 1600 条模型自生成轨迹;无有效采样或 reward 完全相同的 group 会被跳过。

这样的配置减少了每次更新覆盖的题目数量,同时增加了同一道题下的候选回答数量。更大的 group 能为组内平均 reward 和相对 advantage 提供更丰富的比较样本,这正是本次 GRPO 训练信号的来源。

1. 让当前策略对同一道题采样多个答案

每个训练 step 都先从当前 LoRA 权重创建 sampler,再为每道题一次采样 16 个回答:

python 复制代码
sampler = trainer.save_weights_and_get_sampling_client()

result = sampler.sample(
    prompt=trio.ModelInput.from_ints(prompt_tokens),
    num_samples=group_size,
    sampling_params=params,
    return_text=True,
).result()

这里必须使用"当前策略"的权重,因为后续 importance_sampling 需要 rollout 生成时的 old logprobs。同步采样调用返回 future,.result() 表示等待远程采样完成并取得结果。

2. 用数值正确性构造 reward

训练代码不是只给 0 或 1。回答正确且使用 \boxed{} 时 reward 为 1.0;回答正确但没有 boxed 格式时为 0.85;答案错误但数值接近标准答案时,会按相对误差得到较低的 shaping reward,正确格式还能获得少量加分。

python 复制代码
if exact:
    reward_value = 1.0 if boxed else 0.85
elif pred_value is not None and gold_value is not None:
    scale = max(abs(float(gold_value)), 1.0)
    rel_error = abs(float(pred_value - gold_value)) / scale
    reward_value = max(0.0, 0.45 * (1.0 - min(rel_error, 1.0)))
    if boxed:
        reward_value += 0.10
else:
    reward_value = 0.10 if boxed else 0.0

这种设计同时提供"答案是否正确""数值是否接近"和"格式是否合规"三个层次的反馈。不过,最终 79% 与 95% 的准确率只看答案是否与标准值精确相等,不使用 shaping reward 作为准确率。

3. 计算 group-relative advantage

每个 completion 的 advantage 是它的 reward 减去同一道题所有有效 completion 的平均 reward:

python 复制代码
mean_reward = sum(rewards) / len(rewards)
for sample in samples:
    sample["advantage"] = sample["reward"] - mean_reward

同组高于平均分的回答得到正 advantage,低于平均分的回答得到负 advantage。如果一组 reward 的标准差接近 0,代码会跳过该组,避免提交一批全为 0 的训练信号。

4. 对齐 token、old logprobs 与 advantage

GRPO 在 PyTRIO 中使用 importance_sampling loss。prompt token 只提供上下文,不参与训练,因此对应的 target、logprob 和 advantage 都用 0 占位;completion 区间才填入真实值:

python 复制代码
obs_len = len(prompt_tokens) - 1
input_tokens = prompt_tokens + sample["tokens"][:-1]
target_tokens = [0] * obs_len + sample["tokens"]
old_logprobs = [0.0] * obs_len + sample["logprobs"]
advantages = [0.0] * obs_len + [sample["advantage"]] * len(sample["tokens"])

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

input_tokenstarget_tokensold_logprobsadvantages 的长度必须完全一致。这里 obs_len = len(prompt_tokens) - 1,正是为了配合自回归预测时的一位右移。

最后,把有训练信号的 Datum 提交给远程 trainer,并完成一次 Adam 更新:

python 复制代码
fwd = trainer.forward_backward(datums, loss_fn="importance_sampling")
opt = trainer.optim_step(trio.AdamParams(learning_rate=4e-5))
metrics = fwd.result().metrics
opt.result()

训练日志会逐 step 输出平均 reward、精确答对率、组内 reward 标准差、有效 group 数、跳过的同分 group 数、Datum 数量和 loss 指标。训练结束后,脚本会打印可用于推理的 LoRA 权重路径:

text 复制代码
Saved LoRA sampler weights: trio://...

评估结果

复制训练结束时打印的权重路径,然后运行:

bash 复制代码
python eval.py --checkpoint-path 'trio://你的权重路径' --limit 100

eval.py 对 base model 和 checkpoint 使用相同的 test split 前 100 条数据、prompt 模板、答案解析逻辑、temperature=0.0max_tokens=512。它会并发执行采样,但并发只影响评估速度,不改变计分方式。

本次结果如下:

模型 数据范围 采样方式 正确数 Exact Accuracy
Qwen3.5-4B Base GSM8K test 前 100 条 temperature 0 79/100 79%
Qwen3.5-4B GRPO,50 step GSM8K test 前 100 条 temperature 0 95/100 95%

经过 50 step GRPO 训练,准确率从 79% 提升到 95%,绝对提升 16 个百分点。这说明在本次小规模实验中,基于可验证数学答案的 group-relative reward 已经能提供有效的强化学习信号。

从资源消耗来看,GRPO 训练使用了 0.29M Token ,评估阶段使用了 0.06M Token ,训练与评估合计 0.35M Token,总成本为 1.7 元

同时也要注意,这里只评估了 test split 排序后的前 100 条数据,结果来自一次训练运行,并非完整测试集或多随机种子的平均值。因此,它适合用于快速验证 GRPO 流程和训练方向,不应直接当作模型在完整 GSM8K 上的最终成绩。

这次实验最直观的感受是:没想到用不到 2 元,就能通过 RL 让一个 LLM 的准确率提高这么多。 从 79% 到 95% 的结果也让我更直接地感受到,当任务的答案可以被可靠验证时,即使只进行 50 step 的小规模 GRPO 训练,强化学习也可能带来很明显的收益。

常用命令

先评估 base model:

bash 复制代码
python eval.py --limit 100

运行与本文一致的 50 step 训练:

bash 复制代码
python train.py --steps 50 --batch-size 2 --group-size 16 --max-tokens 512 --eval-limit 0

评估训练后的 checkpoint:

bash 复制代码
python eval.py --checkpoint-path 'trio://你的权重路径' --limit 100

如果只想先验证代码链路,可以缩小数据、batch、group 和生成长度:

bash 复制代码
python train.py --limit 8 --steps 2 --batch-size 2 --group-size 2 --max-tokens 128 --eval-limit 0
相关推荐
开开心心_Every1 小时前
电脑文件搜索软件支持内容和拼音搜索
linux·服务器·人工智能·r语言·pdf·音视频·symfony
SLD_Allen1 小时前
从Cloud Native到AI Native:K8s DRA与Agent协议栈
人工智能·云原生·kubernetes·cloud native·ai native
circuitsosk1 小时前
平台整体能力与功能特性设计:分群、联运、ABTest与运营位系统
python·机器学习·搜索引擎·ab测试·vllm·rag检索
wx_xkq12881 小时前
优秘智能:企业AI Agent落地的5大陷阱与工程化解法
人工智能
她说可以呀1 小时前
Spring-ai 2.0 MCP
java·人工智能·spring
2603_954708311 小时前
微能网协调控制箱的核心价值:让多种能源“协同作战”
大数据·运维·网络·人工智能·架构·能源
Java成神之路-1 小时前
Spring AI 核心探秘:四大 Prompt 角色底层设计与完整闭环实战
人工智能·spring·prompt
一次旅行1 小时前
OpenAI 新版提示词指南
人工智能·chatgpt·github
hans汉斯1 小时前
人工智能与机器人研究|面向无标签数据的三维场景语义理解方法研究
人工智能·神经网络·算法·信息可视化·cnn·机器人