尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

REINFORCE算法解析:从策略梯度推导到PyTorch实战与RLHF应用

REINFORCE算法解析:从策略梯度推导到PyTorch实战与RLHF应用 很多学习强化学习的同学第一次见到 RLHFReinforcement Learning from Human Feedback基于人类反馈的强化学习中的 PPO、策略梯度、奖励模型这些概念时都会卡在同一个地方大语言模型生成文本明明是一个离散序列问题为什么突然就用上了强化学习策略梯度里的数学公式到底是怎么推出来的为什么大家都说 REINFORCE 是最基础的策略梯度算法这篇文章不绕弯子直接把 REINFORCE 算法从目标函数到参数更新的每一步数学推导拆开然后用 PyTorch 在 CartPole 环境上跑一个可复现的完整案例最后再回到 RLHF说明大模型训练过程中 REINFORCE 的思想是如何发挥作用的。适合正在入门强化学习或者准备深入研究 RLHF、PPO 的开发者阅读。1. 背景与核心概念1.1 强化学习在解决什么问题强化学习的核心是让智能体Agent通过与环境交互来学习最优策略。智能体在某个状态s下选择动作a环境返回奖励r并进入下一个状态s如此循环往复得到一条轨迹s0, a0, r0, s1, a1, r1, ..., sT智能体要学会的不是“记住某个状态的正确答案”而是学会一种策略π(a|s)也就是在状态s下选择动作a的概率分布。和监督学习不同强化学习没有现成的标签只有一个延迟的、稀疏的奖励信号。比如下围棋要等整盘结束才知道输赢比如训练机械臂抓取物体可能试了很多次才成功一次。这种“延迟奖励”的特点决定了强化学习不能像普通神经网络那样直接做梯度下降而需要一套专门的目标函数和更新方式。1.2 RLHF 与 REINFORCE 的关系RLHF 是让大语言模型对齐人类偏好的关键技术。它一般分为三步先对预训练模型做监督微调SFT让模型学会符合人类表达习惯的回复收集人类对多个回答的偏好排序训练一个奖励模型Reward Model用强化学习优化语言模型让模型在生成回复时获得更高的奖励分数。第 3 步就是强化学习登场的地方。语言模型每生成一个 token可以看成是在状态当前上下文下采取一个动作生成某个词整段生成过程就是一条轨迹。REINFORCE 算法作为最基础的策略梯度方法正好能解释清楚“如何根据最终奖励来调整每一步生成概率”这个核心问题。虽然工业界 RLHF 多数使用 PPO但 PPO 本质上是在 REINFORCE 策略梯度思想上的工程化改进。理解 REINFORCE 的推导是理解 PPO 和 RLHF 的一把钥匙。1.3 策略梯度方法的直观理解在分类任务中我们可以定义交叉熵损失让模型输出的类别概率逼近真实标签。但强化学习没有真实标签只有奖励信号。策略梯度的思路是如果某个动作在事后被证明带来了高回报就增大这个动作的概率如果回报低就减小这个动作的概率。问题在于一次交互中的高回报是多个动作共同作用的结果如何把功劳合理分配到每一步动作上这正是 REINFORCE 数学推导要解决的事情。2. 环境准备与版本说明2.1 开发环境本文示例代码使用 Python 和 PyTorch需要在本地安装以下依赖pip install torch gymnasium matplotlib版本方面不需要完全一致以下环境均可运行Python 3.8 及以上PyTorch 1.13 及以上gymnasium 0.29 及以上注意如果你的环境里安装的还是旧版gymAPI 会有差异建议统一使用gymnasium。本文代码基于gymnasium编写。2.2 为什么选择 CartPole 环境CartPole 是强化学习入门最经典的测试环境。小车在一个一维轨道上移动杆子竖立在小车上。智能体需要左右移动小车让杆子尽量保持竖直。每坚持一个时间步奖励加 1杆子倾斜角度超过 15 度或者小车移出轨道回合结束。这个环境的特点是状态空间是连续的小车位置、速度、杆子角度、角速度动作空间是离散的向左或向右单回合最多 500 步训练速度快适合验证策略梯度类算法的正确性。2.3 项目结构reinforce-demo/ ├── agent.py # 策略网络和基线网络 ├── train.py # 训练入口 └── run.py # 可视化运行演示3. REINFORCE 核心原理与数学公式推导这一节是整个文章的核心。我尽量做到每一步都给出数学解释而不是直接堆公式。3.1 定义目标函数假设策略网络参数为θ给定初始状态分布ρ(s0)智能体在策略π_θ(a|s)下与环境交互。一条轨迹τ的概率可以写成P(τ|θ) ρ(s0) ∏ π_θ(a_t|s_t) P(s_{t1}|s_t, a_t)其中P(s_{t1}|s_t, a_t)是环境的状态转移概率通常未知也不受策略参数θ影响。轨迹的总回报定义为带折扣的累计奖励R(τ) Σ_{t0}^{T-1} γ^t r_t优化目标是找到一组参数θ使得期望回报最大J(θ) E_{τ∼π_θ}[R(τ)]这是一个关于θ的期望函数但我们无法直接对期望求导因为期望里涉及采样过程采样过程不可微。3.2 从期望梯度到对数似然比这里用到强化学习中最经典的技巧叫做 log-derivative trick也叫似然比技巧。先写出梯度的展开形式∇_θ J(θ) ∇_θ ∫ P(τ|θ) R(τ) dτ积分与求导交换顺序后得到∇_θ J(θ) ∫ ∇_θ P(τ|θ) R(τ) dτ关键一步来了利用∇P P ∇log P把梯度转成对数形式∇_θ J(θ) ∫ P(τ|θ) ∇_θ log P(τ|θ) R(τ) dτ发现积分里面恰好是期望形式因此∇_θ J(θ) E_{τ∼π_θ} [ ∇_θ log P(τ|θ) R(τ) ]这是策略梯度定理的第一层结果。它的核心思想是我们不需要对环境求导只需要对策略的对数概率求导然后乘以轨迹回报。环境模型是未知的这件事并不影响我们计算梯度。3.3 轨迹概率分解与化简接着把log P(τ|θ)展开log P(τ|θ) log ρ(s0) Σ_t [ log π_θ(a_t|s_t) log P(s_{t1}|s_t, a_t) ]对θ求梯度时log ρ(s0)和log P(s_{t1}|s_t, a_t)都与策略参数θ无关直接消掉只剩下∇_θ log P(τ|θ) Σ_{t0}^{T-1} ∇_θ log π_θ(a_t|s_t)代入上一节的结果就得到了最终形式的策略梯度∇_θ J(θ) E_{τ∼π_θ} [ Σ_{t0}^{T-1} ∇_θ log π_θ(a_t|s_t) R(τ) ]这个公式的意思是轨迹的每个时间步上策略网络对动作的对数概率求梯度再用整条轨迹的回报作为权重。回报越高梯度方向就越倾向于增大对应动作的概率。3.4 从全轨迹回报到折现回报直接使用整条轨迹R(τ)有一个问题R(τ)包含了当前动作之前所有时间步的奖励。但根据马尔可夫性质t时刻的动作只能影响t时刻之后的奖励不应该为之前的奖励负责。所以更合理的做法是使用 return-to-go也就是从t时刻开始的累计折现回报G_t Σ_{kt}^{T-1} γ^{k-t} r_k由于当k t时E[∇log π_θ(a_t|s_t) r_k]等于 0把R(τ)换成G_t不会改变梯度期望但能明显降低方差。于是策略梯度可以改写为∇_θ J(θ) E_{τ∼π_θ} [ Σ_{t0}^{T-1} ∇_θ log π_θ(a_t|s_t) G_t ]3.5 蒙特卡洛采样与 REINFORCE 更新规则上面的期望无法直接计算只能通过采样来估计。REINFORCE 的做法很直接让智能体在当前策略下跑完一整条轨迹然后用这条轨迹的样本均值来近似梯度。假设采集了N条轨迹策略梯度的经验估计是∇_θ J(θ) ≈ (1/N) Σ_{i1}^N Σ_{t0}^{T_i-1} ∇_θ log π_θ(a_{i,t}|s_{i,t}) G_{i,t}得到梯度后用梯度上升更新参数θ ← θ α ∇_θ J(θ)在实际实现中我们通常把-Σ log π_θ(a_t|s_t) G_t作为损失函数做梯度下降效果等价于上面的梯度上升。3.6 引入基线 Baseline 降低方差REINFORCE 的一个明显问题是方差过大。同一个动作在这次试验中回报高下次试验中回报低梯度方向可能剧烈波动。为了降低方差可以引入一个基线函数b(s_t)把更新式子改成∇_θ J(θ) E_{τ∼π_θ} [ Σ_t ∇_θ log π_θ(a_t|s_t) (G_t - b(s_t)) ]为什么减去基线不影响期望因为E_{a_t∼π_θ} [ ∇_θ log π_θ(a_t|s_t) b(s_t) ] b(s_t) ∇_θ Σ_a π_θ(a|s_t) b(s_t) ∇_θ 1 0所以减去任何只与状态有关的函数期望都不会变但方差可以下降。最简单的基线是回报均值更实用的是学习一个状态价值网络V(s)用G_t - V(s_t)作为带有优势估计的更新信号。3.7 REINFORCE 的局限性REINFORCE 的思想非常优雅但它有几个工程上必须面对的问题方差大单条轨迹的回报波动很大训练不稳定样本效率低每条轨迹用完就丢无法重复利用只能处理回合制任务无法在无限长任务中持续更新训练速度慢需要大量交互次数。后续的 Actor-Critic 算法、GAE广义优势估计、PPO 等技术都是围绕“降低方差、提高样本效率”来改进 REINFORCE 的。4. 完整实战基于 PyTorch 实现 REINFORCE接下来我们把理论变成代码在 CartPole 环境上实现一个带基线的 REINFORCE。4.1 策略网络与基线网络文件路径agent.pyimport torch import torch.nn as nn import torch.nn.functional as F from torch.distributions import Categorical class PolicyNet(nn.Module): 策略网络输入状态输出每个动作的概率 def __init__(self, state_dim, hidden_dim, action_dim): super().__init__() self.fc1 nn.Linear(state_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, action_dim) def forward(self, obs): x F.relu(self.fc1(obs)) x self.fc2(x) return F.softmax(x, dim-1) def get_action(self, obs): 根据当前状态采样动作同时返回动作的对数概率 obs torch.as_tensor(obs, dtypetorch.float32) probs self.forward(obs) dist Categorical(probs) action dist.sample() return action.item(), dist.log_prob(action) class ValueNet(nn.Module): 价值网络输入状态输出状态价值估计 V(s)作为 baseline def __init__(self, state_dim, hidden_dim): super().__init__() self.net nn.Sequential( nn.Linear(state_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, 1) ) def forward(self, obs): return self.net(obs).squeeze(-1)这里有一个容易忽略的细节策略网络的输出是动作概率分布所以用了Categorical来采样动作并保留对数概率。反向传播时模型参数会被记录在计算图上损失函数会通过采样的动作路径回传梯度。4.2 实现训练流程文件路径train.pyimport gymnasium as gym import numpy as np import torch import torch.nn as nn import torch.optim as optim from agent import PolicyNet, ValueNet def compute_returns(rewards, gamma0.99): 计算每个时间步的折现回报 G_t returns [] G 0 for r in reversed(rewards): G r gamma * G returns.insert(0, G) return torch.tensor(returns, dtypetorch.float32) def reinforce(env, policy_net, value_net, policy_opt, value_opt, gamma0.99, num_episodes500, entropy_coef0.01): episode_rewards [] for episode in range(num_episodes): obs, _ env.reset() log_probs [] rewards [] states [] while True: action, log_prob policy_net.get_action(obs) next_obs, reward, terminated, truncated, _ env.step(action) done terminated or truncated log_probs.append(log_prob) rewards.append(reward) states.append(obs) obs next_obs if done: break # 计算折现回报 returns compute_returns(rewards, gamma) # 标准化回报等价于使用常数基线能稳定训练 returns (returns - returns.mean()) / (returns.std() 1e-9) # 计算状态价值基线 states_tensor torch.as_tensor(np.array(states), dtypetorch.float32) values value_net(states_tensor) # 优势 回报 - 基线 advantages returns - values.detach() # 策略损失最大化对数概率 * 优势 log_probs_tensor torch.stack(log_probs) policy_loss -(log_probs_tensor * advantages).mean() # 熵正则鼓励探索防止过早收敛到局部最优 obs_tensor torch.as_tensor(np.array(states), dtypetorch.float32) probs policy_net(obs_tensor) dist torch.distributions.Categorical(probs) entropy dist.entropy().mean() policy_loss policy_loss - entropy_coef * entropy # 价值损失让 V(s) 更接近真实回报 value_loss nn.MSELoss()(values, returns) # 更新 policy_opt.zero_grad() policy_loss.backward() policy_opt.step() value_opt.zero_grad() value_loss.backward() value_opt.step() episode_rewards.append(sum(rewards)) if (episode 1) % 50 0: avg_reward np.mean(episode_rewards[-50:]) print(fEpisode {episode 1}, Average Reward: {avg_reward:.2f}) return episode_rewards if __name__ __main__: env gym.make(CartPole-v1) state_dim env.observation_space.shape[0] action_dim env.action_space.n hidden_dim 128 policy_net PolicyNet(state_dim, hidden_dim, action_dim) value_net ValueNet(state_dim, hidden_dim) policy_opt optim.Adam(policy_net.parameters(), lr1e-3) value_opt optim.Adam(value_net.parameters(), lr1e-2) rewards reinforce(env, policy_net, value_net, policy_opt, value_opt)这段代码有几个值得说明的地方。第一returns做了标准化。标准化相当于把均值作为基线在训练初期能明显提升稳定性。虽然理论上的基线是状态相关的V(s)但常数基线同样不影响期望只是一个比较粗略的降方差手段。第二values.detach()是为了避免优势计算时把梯度传到价值网络。价值网络有自己的损失函数路径必须分离。第三熵正则项不是 REINFORCE 的必要部分但在实践中几乎必加。它鼓励策略保持一定的随机性防止模型在探索早期就“锁死”在某个子优动作上。4.3 运行与验证在项目根目录执行python train.py预期会看到类似下面的输出Episode 50, Average Reward: 23.78 Episode 100, Average Reward: 42.15 Episode 150, Average Reward: 71.53 ... Episode 500, Average Reward: 178.24这里的数字只是演示实际数值会因随机种子和超参数不同而波动。注意 CartPole 单回合上限 500 步如果平均奖励稳定在 400 以上说明策略已经基本学好了。4.4 可视化训练曲线REINFORCE 的训练曲线通常抖动非常明显这是该算法的正常表现。想要更直观地观察收敛过程可以记录每个 episode 的奖励并绘制曲线。在train.py中加入import matplotlib.pyplot as plt rewards reinforce(env, policy_net, value_net, policy_opt, value_opt) plt.plot(rewards) plt.xlabel(Episode) plt.ylabel(Total Reward) plt.title(REINFORCE on CartPole-v1) plt.show()曲线整体会呈现上升趋势但中间会有很多起伏。如果完全看不到上升趋势优先检查学习率、折扣因子和熵系数。5. 从 REINFORCE 到 RLHF大模型策略优化这一节把视角从 CartPole 拉回大模型训练看看 REINFORCE 的思想如何自然延伸到 RLHF 中。5.1 大模型生成过程就是强化学习轨迹要理解 RLHF首先要接受一个观点语言模型的生成过程可以看成强化学习中的一次轨迹交互。给定一个提示词x语言模型逐个生成 tokeny1, y2, ..., yT。在这个过程中状态s_t是提示词和已经生成的 token 序列也就是当前上下文动作a_t是下一个要生成的 token策略π_θ就是语言模型本身环境是语言模型内部的解码过程它的状态转移是确定性的奖励r_T在整段文本生成结束后由奖励模型给出。于是语言模型 RLHF 的策略优化目标可以写成max_θ E_{x∼D, y∼π_θ(·|x)} [ R_φ(x, y) ]其中R_φ(x, y)是奖励模型对整段回复的打分。5.2 奖励模型是怎么训练的奖励模型通常基于一个 SFT 模型在最后一层加一个线性层输出一个标量分数。训练数据是人类对多个回答的偏好比较。假设对于同一个提示x人类认为回答y_w比回答y_l更好奖励模型的训练使用 Bradley-Terry 模型P(y_w y_l) σ(R_φ(x, y_w) - R_φ(x, y_l))其中σ是 sigmoid 函数。训练目标是最大化人类偏好被预测正确的概率。这个步骤和强化学习本身没有直接关系但必须完成它后续的策略优化才有奖励信号来源。5.3 REINFORCE 视角下的 RLHF 更新用 REINFORCE 的思想来看优化目标对模型参数θ的梯度是∇_θ E[R(x, y)] E[ Σ_t ∇_θ log π_θ(y_t | x, y_t) R(x, y) ]这个公式和 CartPole 上的策略梯度公式结构完全一致。每个 token 的对数概率梯度乘以整段回复的奖励分数。如果某段回复获得了高分模型会加大这一段所有 token 的生成概率。但在实际 RLHF 中直接使用 REINFORCE 会有严重问题。语言模型的输出空间巨大单次采样方差极高训练很不稳定。所以工业界几乎不会直接用 REINFORCE而是在它基础上引入 PPO 的裁剪目标、GAE 优势估计、重要性采样等机制。5.4 为什么 PPO 会成为 RLHF 的标准选择PPO 本质上是解决“REINFORCE 方差大、样本效率低”这两个问题的工程化方案。它的核心改进有四点重要性采样允许模型利用旧策略采样的数据来更新新策略提高样本复用率Clip 目标限制策略更新的幅度防止一步更新过大导致策略崩塌优势估计用 GAE 代替单步回报在偏差和方差之间取得平衡KL 惩罚在奖励中加入与参考模型的 KL 散度惩罚防止模型偏离人类语言分布太远。所以理解 RLHF 时不能只记住 PPO 这个名词而是要看到它背后的策略梯度本质——也就是 REINFORCE 推导出来的那个式子。5.5 简化版 REINFORCE 式 RLHF 伪代码为了帮助理解下面给出一个简化版的 REINFORCE 式 RLHF 更新伪代码。它不能直接用于生产环境但能清晰表达核心逻辑for each step: # 1. 采样 x sample_prompt_from_dataset() y generate_response(x, policy_model) # 2. 计算奖励 reward reward_model(x, y) - β * kl_divergence(policy_model(y|x), ref_model(y|x)) # 3. 计算策略梯度损失 log_probs sum(log policy_model(y_t | x, y_t) for t in range(len(y))) loss - log_probs * reward # 4. 更新策略模型 optimizer.zero_grad() loss.backward() optimizer.step()这里唯一的区别是大模型生成的一次完整回复被当成一条轨迹每个 token 的生成概率都乘上同一个奖励。实际训练时为了防止模型为了高分强行改变语言风格需要在奖励中减去 KL 散度惩罚这也是 RLHF 实践中非常重要的细节。6. 常见问题与排查思路REINFORCE 代码写起来很简单调试起来却不容易。下面列出我实际训练过程中遇到频率最高的问题。问题现象常见原因解决思路训练很久奖励不上升学习率过大或过小把学习率先调到 1e-3 ~ 3e-3观察前 100 个 episode 曲线奖励曲线波动非常大REINFORCE 方差大的固有特性增加批量轨迹数量引入价值网络作为 baseline策略过早收敛到单一动作熵系数缺失或过小添加熵正则熵系数一般取 0.01 ~ 0.1梯度出现 NaN回报数值过大对 returns 做标准化或降低学习率训练开始后很快达到 500 分但随后崩塌策略更新幅度过大减小学习率或者改用 PPO 的 clip 机制gym 环境报错reset 返回值数量不对gym 与 gymnasium 版本混用统一使用 gymnasium注意新版 reset 返回 (obs, info)loss 一直为负高回报轨迹导致负对数似然乘正数这是正常现象关注 total reward 而非 loss单条轨迹更新特别不稳定只用一个 episode 估计梯度每次更新前跑多个 episode把梯度累计起来再更新三个最常见的调试建议第一先跑通最简单的版本。不要一开始就加各种技巧先用原始 REINFORCE 在 CartPole 上跑通确认策略梯度方向是对的再逐步加 baseline、熵正则、批量更新。第二监控多组指标。除了总奖励还应该记录平均熵、V(s) 的 loss、梯度范数。熵降为零说明策略完全确定很可能停止探索梯度范数突变往往是学习率过大。第三固定随机种子。训练强化学习代码前先固定torch.manual_seed和np.random.seed这能让你在调试时确认某个修改到底有没有效果。7. 最佳实践与工程建议7.1 理解优先于调参REINFORCE 的超参数不多但每个都很敏感。有人喜欢疯狂调学习率却不理解回报计算和基线的作用。实际调试时先把公式和代码逐行对应上比盲目调参有效得多。比如G_t的计算方向、detach()的位置、熵正则的符号任何一个写错都会让训练行为怪异。7.2 用批量轨迹降低方差每次收集一条轨迹就更新是最原始的 REINFORCE 写法但方差非常大。工程上更推荐的做法是一次性采样多个 episode把它们拼接成一整个 batch计算所有轨迹的 returns 并做标准化用同一个 batch 的数据更新一次参数重复上述过程。批量更新能让梯度方向更稳定训练曲线更平滑。7.3 价值网络与策略网络分开代码里把策略网络和价值网络分成两个独立的网络原因在于两者学习目标不同策略网络需要最大化期望回报价值网络需要拟合状态价值函数。如果共享网络结构策略梯度和价值损失会相互干扰。即使是小型示例也建议分开建模。这个习惯在后续学习 Actor-Critic 时会很有帮助。7.4 日志、种子与可复现性强化学习实验很难复现因为环境本身的随机性很大。在生产实验代码中至少要做到固定所有随机种子保存每个 episode 的奖励、熵、损失到 CSV 或 TensorBoard保存模型权重时同时保存超参数配置每次实验记录环境版本和依赖版本。这些看起来麻烦但在做算法调优时能节省大量时间。7.5 在真实场景迁移到 RLHF 的注意点如果把 REINFORCE 的思想迁移到 RLHF 或机器人控制场景有几个额外的风险需要注意。第一是奖励设计。奖励模型打分可能被模型“钻空子”生成高分但不符合人类意图的内容。解决办法是加入 KL 惩罚并定期用人工评估校准奖励模型。第二是安全边界。强化学习策略在上线前必须在仿真环境或沙盒中充分验证。无论是语言模型还是机械臂控制都不能直接让未经验证的策略接触真实用户或真实设备。合法授权、最小权限、灰度发布这些原则同样适用于 RLHF 模型的部署。第三是样本效率。语言模型每次采样成本很高不能像 CartPole 那样轻松跑几十万次。这也是 PPO 这类高样本效率算法成为主流的原因之一。如果数据来源是历史日志可以关注离线强化学习方向比如 IQLImplicit Q-Learning它利用固定数据集学习策略适合无法在线交互的训练场景。8. 总结与后续学习建议读到这里你已经走完了从 REINFORCE 理论推导到代码实现再到 RLHF 应用分析的完整链路。需要掌握的核心内容可以浓缩成几句话策略梯度的核心公式通过 log-derivative trick 得到不需要对未知的环境模型求导轨迹概率分解后梯度只依赖策略网络的对数概率和回报用 return-to-go 代替整条轨迹回报用 baseline 降低方差是 REINFORCE 落地的关键RLHF 中的语言模型生成过程本质上是一条 token 级轨迹PPO 是在 REINFORCE 基础上的方差降低与样本效率改进。下一步的学习路线我建议顺着这条路径走在 CartPole 上把 REINFORCE 换成带基线的版本体验方差下降带来的稳定性提升学习 Actor-Critic 算法理解价值网络如何参与策略更新深入 GAE 原理弄明白“优势估计”到底比单步回报好在哪里再看 PPO 的 clip 目标回到 RLHF 的 PPO 训练流程大部分公式就能对上号了如果你对连续控制感兴趣可以再做机械臂仿真环境比如 MuJoCo、PyBullet上的强化学习实战方向上会接触到 DDPG、TD3、SAC 这类算法。最后给你的实践建议是不要只盯着最新算法REINFORCE 虽然简单但它把“通过奖励调整动作概率”这个核心思想讲得足够清楚。真正把 REINFORCE 的公式和代码一一对应之后再去看 PPO、RLHF 的论文和开源实现思路会顺畅很多。
返回列表