TRL与VERL中的GRPO算法实现

GRPO算法最早在DeepSeek-Math论文中提出,通过多次采样获取相对优势取消了Critic模型,大幅降低了训练复杂度。

首先回顾总体GRPO算法流程

其中整体优化公式如下:

$$
\mathcal{L}{GRPO} = -\frac{1}{G}\sum{i=1}^G \frac{1}{|o_i|}\sum_{t=1}^{|o_i|} \min\left[\underbrace{\frac{\pi_\theta(o_{i,t}|q,o_{i,<t})}{\pi_{\theta_{old}}(o_{i,t}|q,o_{i,<t})}}{\text{coef_1}} \hat{A}{i,t},\ \text{clip}\left(\frac{\pi_\theta}{\pi_{\theta_{old}}}, 1{-}\epsilon, 1{+}\epsilon\right) \hat{A}{i,t}\right] + \beta, \mathbb{D}{KL}[\pi_\theta | \pi_{ref}]
$$

其中,
$$\hat{A}{i,t} = \dfrac{r_i - \text{mean}({r_1,…,r_G})}{\text{std}({r_1,…,r_G})}$$
$$\mathbb{D}
{KL}[\pi_\theta | \pi_{ref}] = \sum_{a} \pi_\theta(a|s) \log \frac{\pi_\theta(a|s)}{\pi_{ref}(a|s)}={\frac{\pi_{ref}(o_i|q)}{\pi_\theta(o_i|q)}}-\log{\frac{\pi_{ref}(o_i|q)}{\pi_\theta(o_i|q)}}-1$$

整体的训练流程大致如下:

  1. 加载数据集即prompts
  2. 执行rollout过程,针对每个prompts生成G个completions,和old_logits
  3. 计算每个completions的reward和advantage
  4. 拼接prompt + completions进入policy model得到 logits
  5. prompt+completions进入ref_model得到ref logits,计算KL散度
  6. 根据GRPO公式计算损失并进行反向传播更新模型参数

TRL中GRPO算法的实现

TRL是huggingface官方发布的用于LLM RL训练的组件库,内置了多种基于trasformers架构模型的RL算法,本文重点介绍GRPO算法。基于commit 4890abf0ac0940dc23eeaa6e438d6e4e2217260d

其训练流程实现文件在grpo_trainer.py中。配置文件在grpo_config.py中。

TRL的训练循环依赖于transformers库,在transformers/trainer.py中的training_step()函数中实现。流程为

1
_prepare_inputs(inputs) -> compute_loss() -> backward()

在TRL的GRPOTrainer中,重写了_prepare_inputs()函数,后续调用链为:

1
2
3
4
5
6
7
8
9
10
_prepare_inputs(generation_batch)
|_> _generate_and_score_completions(generation_batch) -> output{prompt_ids, completion_ids, logprobs,advantages}
|_> _get_per_token_logps_and_entropies() -> old_per_token_logps, ref_per_token_logps
|_> _calculate_rewards(inputs,..) -> rewards_per_func
|_> advanatages



compute_loss()
|_> _compute_loss()

因此重要的逻辑在_generate_and_score_completions()函数和_compute_loss()函数中。

数据加载

get_train_dataloader()函数中,TRL使用自定义的方式加载数据集,为了避免每次训练都需要一次数据生成的低效,这里首先取一个很大的batch,也即代码中的

1
batch_size=self._train_batch_size * self.args.steps_per_generation,  # < this is the change

steps_per_generation可以理解为首先取很多数据,再分成steps_per_generation个小batch进行训练。

num_iterations 是指一次采集的样本用于多少次训练,当该值大于1时,需要进行重要性采样的纠偏。

采样生成

即函数_generate_and_score_completions()

接收输入inputs,其为list[dict]类型,每个dict为一条训练样本,长度为per_device_train_batch_size * steps_per_generation

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
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
┌─────────────────────────────────────────────────────────────┐
│ DataLoader 层 │
│ get_train_dataloader() │
│ batch_size = per_device_train_batch_size × steps_per_gen │
│ │
│ RepeatSampler (在 _get_train_sampler 中) │
│ 每个 prompt 重复 num_generations=G 次 │
│ → [q0,q0,q1,q1,...] (G=2时) │
└──────────────────────────┬──────────────────────────────────┘
│ generation_batch: [per_device_train_batch_size]

┌─────────────────────────────────────────────────────────────┐
│ _prepare_inputs() │
│ │
│ 每 steps_per_generation × num_iterations 步才重新生成一次 │
│ 否则复用 _buffered_inputs 中缓存的 batch │
└──────────────────────────┬──────────────────────────────────┘
│ 触发生成

┌─────────────────────────────────────────────────────────────┐
│ _generate_and_score_completions() │
│ │
│ ① _generate(prompts) │
│ _tokenize_prompts() → prompt_ids │
│ _generate_single_turn() │
│ ├─ vLLM 路径: VLLMGeneration.generate() │
│ └─ transformers 路径: model.generate() │
│ → prompt_ids_list, completion_ids_list, logprobs_list │
│ │
│ ② pad + to_tensor │
│ prompt_ids: (B, P) padding_side="left" │
│ completion_ids: (B, C) padding_side="right" │
│ attention_mask: (B, P+C) │
│ │
│ ③ 计算 old_per_token_logps (π_θ_old) ←── 旧策略log概率 │
│ _get_per_token_logps_and_entropies(self.model, ...) │
│ shape: (B, C) │
│ (仅在 steps misaligned 或 vLLM 时计算,否则 None) │
│ │
│ ④ 计算 ref_per_token_logps (π_ref) ←── 参考模型log概率 │
│ _get_per_token_logps_and_entropies(self.ref_model, ...) │
│ shape: (B, C) (仅 beta != 0 时) │
└──────────────────────────┬──────────────────────────────────┘


┌─────────────────────────────────────────────────────────────┐
│ _calculate_rewards() │
│ │
│ 对每个 reward_func: │
│ ├─ nn.Module: 调用分类模型 → logits[:,0] │
│ ├─ sync callable: reward_func(prompts, completions, ...) │
│ └─ async callable: asyncio.gather 并发执行 │
│ │
│ rewards_per_func: (B×G_all_procs, num_funcs) │
│ accelerator.gather() 跨进程聚合 ← 因为要组内归一化 │
└──────────────────────────┬──────────────────────────────────┘
│ rewards_per_func: (B_global, F)

┌─────────────────────────────────────────────────────────────┐
│ Advantage 计算 (组内相对归一化) │
│ │
│ rewards = sum(rewards_per_func × weights, dim=1) │
│ shape: (B_global,) = (num_prompts × G,) │
│ │
│ reshape → (num_prompts, G) │
│ mean_grouped = rewards.mean(dim=1) → repeat_interleave(G) │
│ std_grouped = rewards.std(dim=1) → repeat_interleave(G) │
│ │
│ advantages = (rewards - mean_grouped) / (std_grouped + ε) │
│ ↑ 这就是论文中的 Â_{i,t} (此时还不含 token 维度) │
│ │
│ process_slice = [proc_idx*B : (proc_idx+1)*B] │
│ advantages = advantages[process_slice] ← 取本进程的 slice │
└──────────────────────────┬──────────────────────────────────┘
│ inputs dict with advantages, ids, masks,
│ old_per_token_logps, ref_per_token_logps

┌─────────────────────────────────────────────────────────────┐
│ compute_loss() → _compute_loss() (每个 gradient step) │
│ │
│ ① per_token_logps, entropies │
│ = _get_per_token_logps_and_entropies(model, ...) │
│ → log π_θ(o_t | q, o_<t), shape: (B, C) │
│ │
│ ② old_per_token_logps = inputs["old_per_token_logps"] │
│ or per_token_logps.detach() (若 None) │
│ │
│ ③ log_ratio = per_token_logps - old_per_token_logps │
│ = log(π_θ / π_θ_old), shape: (B, C) │
│ │
│ ④ coef_1 = exp(log_ratio) = π_θ / π_θ_old │
│ coef_2 = clamp(coef_1, 1-ε_low, 1+ε_high) ← PPO clip │
│ │
│ ⑤ GRPO loss (loss_type="grpo"): │
│ per_token_loss1 = coef_1 × advantages │
│ per_token_loss2 = coef_2 × advantages │
│ per_token_loss = -min(loss1, loss2) │
│ │
│ ⑥ KL 惩罚 (beta != 0): │
│ per_token_kl = exp(log_ref - log_θ) │
│ - (log_ref - log_θ) - 1 │
│ (Schulman 近似的 KL: exp(x)−x−1 ≥ 0) │
│ per_token_loss += beta × per_token_kl │
│ │
│ ⑦ 归约: │
│ loss = mean_over_batch( │
│ sum_over_tokens(per_token_loss × mask) │
│ / sum(mask).clamp(1) │
│ ) / gradient_accumulation_steps │
└─────────────────────────────────────────────────────────────┘

三、关键设计细节

1. num_generations (G) 怎么流入数据

grpo_trainer.py 中 RepeatSamplermini_repeat_count=G,让同一个 prompt 的 G 份 token ID 连续出现,保证同 GPU 上的 G 个 completion 来自同一个 prompt,便于组内归一化。

2. old_per_token_logps 何时为 None

1
2
3
4
5
6
# _generate_and_score_completions 中
generate_every = steps_per_generation × num_iterations
if gradient_accumulation_steps % generate_every != 0 or use_vllm:
old_per_token_logps = computed # 需要 IS 修正
else:
old_per_token_logps = None # 直接用 detach(),省一次前向

num_iterations=1, steps_per_gen ≤ grad_accum 时,生成和更新步对齐,π_θ_old ≡ π_θ.detach(),无需额外计算。

3. _get_per_token_logps_and_entropies 的 token 维度对齐

1
2
3
4
5
logits = model(...).logits          # (B, P+C, V)
logits = logits[:, :-1, :] # shift left (predict next token)
logits = logits[:, -logits_to_keep:, :] # 只保留 completion 部分
logits /= temperature
logps = selective_log_softmax(logits, completion_ids) # (B, C)

这是标准的 causal LM teacher-forcing log-prob 提取,只对 completion token 计算,避免对 prompt 计算不必要的 logit。

4. KL 散度的近似形式

代码使用 $\text{KL}(π_\theta | π_{ref}) \approx e^{r-1} - r + 1 - 1$(其中 $r = \log π_{ref} - \log π_\theta$),即 Schulman 的无偏近似,保证 KL ≥ 0。

5. 多种 loss_type

loss_type 公式特点 代码实现
grpo 标准 PPO clip,序列平均 min(coef_1, clamp) × A
bnpo token 全局平均(无序列归一化) 分子 sum / 全局 mask sum
dr_grpo 除以 B × max_len(固定分母) DAPO 论文变体
dapo num_items_in_batch 归一化 动态 token 数归一化
cispo 单边 clip(只限 coef_1 上界) clamp(coef_1, max=ε_high)
vespo Gamma 权重替代 PPO clip φ(w) × A × log_prob

四、一次完整训练迭代的时序

1
2
3
4
5
6
7
8
9
step 0  ──► 生成 G×B completions(昂贵)
↓ 计算 rewards & advantages
↓ 缓存到 _buffered_inputs[0..steps_per_gen]
↓ 取 _buffered_inputs[0]compute_loss() → backward()

step 1 ──► 复用缓存,取 _buffered_inputs[1]compute_loss() → backward()
...
step S-1 → 复用缓存,取 _buffered_inputs[S-1] → optimizer.step()
step S ──► 重新生成(num_iterations=1时)

steps_per_generation 控制每次生成的样本被复用多少次(即 μ 步),num_iterations 控制每个生成批次做多少轮完整的梯度累积循环,对应论文中的 μ(micro-iterations)

VERL中的GRPO算法实现

verl即Volcano Engine Reinforcement Learning,是一种高效的分布式RL算法训练引擎,在多个节点之间执行单控制器范式,节点内执行多控制器范式,设计了一套分层API,解耦并封装了复杂RLHF数据流中的计算和数据依赖关系,允许高效的操作编排,并且包含了一个3D混合引擎,在训练和生成之间实现高效的actor模型的重分片。

verl使用ray引擎进行分布式训练,使用GRPO算法进行训练时,启动入口在verl/trainer/main_ppo.py。首先是main()函数读取配置信息,run_ppo()启动ray集群,

  1. 启动TaskRunner实例使用runner.run.remote(config)启动训练,阻塞在ray.get()处,等待训练结束
  2. 启动TaskRunnerrun方法注册actor,critic,reward_model等到集群中,创建数据集
  3. 创建RayPPOTrainer实例,启动fit方法开始训练。

此后的流程遵循GRPO一贯的运算流程,rollout,advantage计算, 及compute loss等

verl中的关键参数

  1. train_batch_size: 即一次性取多少prompt做训练
  2. ppo_mini_batch_size: 每个prompt经过rollout后得到G个completion,ppo_mini_batch_size即指一次性取多少completion进入后续的训练。
  3. ppo_micro_batch_size_per_gpu: 每个GPU上微批次的大小,用于梯度累积。ppo_mini_batch_size个completion可能无法一次性在GPU上训练,将其划分为micro_batch为一组的小batch进行梯度累积。

在rollout model和actor model分离的情形下,一般涉及到四个logprobs。

  1. old_log_probs: 即旧策略模型的log概率,该值每个rollout batch计算一次,也即是在一个train_batch_size里使用actor model计算一次,然后缓存下来。
  2. rollout_log_probs: 即rollout模型的log概率,该值每个rollout batch计算一次,也即是在一个train_batch_size里使用actor model计算一次,然后缓存下来。由于执行rollout的模型的运行参数(如使用vLLM)可能与actor model不同,所以需要进行$\pi_{current}/\pi_{rollout}$的修正。
  3. current_log_probs: 即当前策略模型的log概率,该值在每个micro_batch前向计算时得到。
  4. ref_log_probs: 即参考策略模型的log概率,该值用于计算KL散度,通常在每个rollout batch计算一次。

重要性采样的计算有两类,一次是弥合rollout模型和actor模型运行环境的不同,使用$\pi_{current}/\pi_{rollout}$计算,一次是弥合当前策略模型和旧策略模型的不同,使用$\pi_{current}/\pi_{old}$计算。


TRL与VERL中的GRPO算法实现
https://wenzhaoabc.github.io/llm/trl_verl/
作者
wenzhaoabc
发布于
2025年10月1日
许可协议