GRPOTrainer 详解:可验证奖励与组相对优化

05-01 偏好对齐选型 中,GRPO(Group Relative Policy Optimization,组相对策略优化)适合「答案可自动判定对错」的任务——代码单测、数学 \boxed{} 校验、格式规则等。本篇介绍 TRL GRPOTrainer 的工程用法;与 DPO 对比见 06-01 DPOTrainer

段末注释:GRPO 对同一 prompt 采样一组 completion,用组内 reward 的相对高低更新策略,无需单独 critic 网络;DeepSeek-R1 等推理模型训练路径的代表方法之一。

系列索引:微调技术路线导读


一、在流水线中的位置

1
2
3
4
5
6
7
SFT checkpoint(会格式、会推理骨架)

├─ 定义 reward_fn(completions, prompts) → 标量

└─ GRPOTrainer:每组 G 个采样 → 组内相对优势 → 更新 LoRA

└─ 评估 pass@k / 准确率 → [04-10 pass@k](./04-评估指标-10-pass@k.md)

与 DPO 的分工

DPO GRPO
信号 人类/模型偏好对 可执行 reward
数据 chosen / rejected prompt + reward 函数
典型任务 对话、安全性 代码、数学、结构化输出

二、核心机制(直觉)

对 prompt (x),采样 (G) 个输出 ({y_1,\ldots,y_G}),得到 reward ({r_1,\ldots,r_G})。GRPO 用组内 baseline(如均值)构造相对优势,拉高高于组均的 completion 概率、压低低于组均的——类似 REINFORCE 的去 baseline 版本,但不需要 value network。

关键:reward 必须可自动计算且与业务目标一致;错误 reward 会强化幻觉或错误格式。


三、依赖与最小示例

1
pip install "trl>=0.16" transformers peft datasets accelerate
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
import re
from datasets import Dataset
from trl import GRPOConfig, GRPOTrainer
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig

MODEL_ID = "./lora-sft-or-base"
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)

# 示例:数学答案在 \boxed{} 中,与 ground_truth 字符串比对
def reward_fn(completions, ground_truth, **kwargs):
rewards = []
for completion, gt in zip(completions, ground_truth):
m = re.search(r"\\boxed\{([^}]+)\}", completion)
pred = m.group(1).strip() if m else ""
rewards.append(1.0 if pred == gt.strip() else 0.0)
return rewards

train_ds = Dataset.from_dict({
"prompt": ["Solve: 2+2=? Put answer in \\boxed{}."],
"ground_truth": ["4"],
})

lora_config = LoraConfig(r=16, lora_alpha=32, task_type="CAUSAL_LM", target_modules="all-linear")

training_args = GRPOConfig(
output_dir="./grpo-out",
num_generations=4, # 每组采样数 G
per_device_train_batch_size=1,
gradient_accumulation_steps=4,
learning_rate=1e-5,
num_train_epochs=1,
max_completion_length=256,
bf16=True,
gradient_checkpointing=True,
)

trainer = GRPOTrainer(
model=MODEL_ID,
reward_funcs=reward_fn,
args=training_args,
train_dataset=train_ds,
peft_config=lora_config,
processing_class=tokenizer,
)
trainer.train()

实际字段名(如 ground_truth 如何传入 reward_fn)以当前 TRL 版本文档为准;上例展示「prompt 列 + reward 函数读 completion」模式。


四、reward 函数设计

任务 reward 示例
代码 单元测试 pass → 1,否则 0(pass@k
数学 \boxed{} 内答案 == 标准解
JSON 格式 json.loads 成功 + schema 校验
分类 解析标签 ∈ 允许集 → 1

实践建议

  1. 先在 100 条上统计 reward 分布;若全 0 或全 1,训练无效。
  2. 可加塑形 reward(部分分、编译通过但未过测给 0.3),但保持简单可解释。
  3. 01-02 清洗 一致:ground_truth 规范统一。

代码任务 reward 骨架

1
2
3
4
5
6
7
8
9
def code_reward(completions, test_cases, **kwargs):
rewards = []
for code, tests in zip(completions, test_cases):
try:
ok = run_tests(extract_code(code), tests) # 沙箱执行
rewards.append(1.0 if ok else 0.0)
except Exception:
rewards.append(0.0)
return rewards

安全:必须在隔离沙箱中执行模型生成代码。


五、GRPOConfig 关键参数

参数 典型值 说明
num_generations 4–16 组大小 (G);越大梯度越稳,越慢
learning_rate 1e-6–1e-5 通常低于 SFT
max_completion_length 256–2048 推理链长任务需更大
temperature(采样) 0.7–1.0 探索多样性
beta / KL 项 依 TRL 版本 约束离 SFT 过远

显存:每组需 G 次生成 + 反传,开销大于 DPO。叠加 07-01 显存优化gradient_checkpointing、小 batch、LoRA。


六、与 SFT 的衔接

步骤 说明
1 SFT 教会 \boxed{} / code block 等输出格式
2 GRPO 在格式之上优化 reward
3 保留 SFT adapter 备份;GRPO 后对比 08-01 基线

无 SFT 直接 GRPO:模型可能不跟格式,reward 长期为 0。


七、评估

指标 用途
reward mean 训练监控
pass@1 / pass@k 代码主指标
数学准确率 抽测 hold-out
通用能力抽检 遗忘

八、常见踩坑

现象 对策
reward 全 0 检查 SFT 格式;放宽 parse;看生成样例
训练极慢 num_generations;缩短 completion
过拟合单测 增 test 多样性;held-out 集
推理链胡编 reward 只认最终答案会鼓励猜;可加步骤校验

九、小结

GRPOTrainer = SFT 格式 + 可验证 reward + 组内相对优化。对话偏好用 06-01 DPO;代码/数学/可校验任务用 GRPO。

下一步:08-02 遗忘与版本管理 | 选型 05-01

-------------本文结束感谢您的阅读-------------