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

资讯详情

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

【强化学习】Hands-on Modern RL项目实践|PPO 算法原理与 PyTorch 实现详解

【强化学习】Hands-on Modern RL项目实践|PPO 算法原理与 PyTorch 实现详解 本文承接上一篇参考开源课程 Hands-On Modern RL 第 1.1 节的讲解框架并结合其配套代码 2-pytorch_ppo.py把 Stable-Baselines3 里那一行model.learn()拆开看看 PPO 算法内部到底在做什么。一、从跑通训练到看懂训练如果你用过 Stable-Baselines3SB3大概率写过这样的代码modelPPO(MlpPolicy,env,verbose1)model.learn(total_timesteps80000)短短一行model.learn()几秒钟内就能让 CartPole 小车从随机乱动学会稳稳地把杆子立住、拿满 500 分。但这行代码背后到底发生了什么如果不搞清楚后面学策略梯度、Actor-Critic、GRPO 这些内容时PPO 就会一直是一个凭空出现的黑盒。这篇文章要做的事情就是把这个黑盒拆开用不到 200 行纯 PyTorch 代码重新实现一遍 SB3 内部发生的事情。整个实现可以拆成三个部分一个做决策的 Actor-Critic 网络、一段收集经验数据的 Rollout 循环、一套根据回报信号调整参数的 PPO 更新规则。掌握了这三部分你就理解了model.learn()的全部本质。在动手之前先快速回顾一下 CartPole 环境本身的设定这是后面所有代码的基础。二、CartPole 环境状态、动作、奖励CartPole 是强化学习里最经典的入门任务一根杆子通过关节连接在一辆小车上小车只能左右移动目标是让杆子尽可能长时间保持竖直不倒。状态观测环境每一帧都会给智能体一个 4 维向量分别对应小车位置、小车速度、杆子角度、杆子角速度。其中速度和角速度理论上没有硬性上限物理引擎每帧现算但实际训练中大多落在 -3~3 之间小车位置的边界是 ±2.4杆子角度的边界约为 ±12°一旦超出这两个阈值回合立刻结束。动作只有两个选项——向左推小车0或向右推小车1没有轻推不推这种中间状态是最简单的离散动作空间。奖励规则极其朴素——每存活一步包括结束的那一步就获得 1 分直到杆子倒下、小车出界或者步数达到 500 步的上限也就是满分。这个奖励设计背后有一层值得注意的含义智能体只知道这一局总共活了多少步却不知道具体是哪一步的决策导致了提前结束。这正是强化学习和监督学习的本质区别——没有人告诉智能体第 23 步应该往左推它只能通过大量试错慢慢琢磨出哪些动作更有利于长期生存。把状态、动作、奖励串起来还差最后一个要素——策略Policy一个从状态到动作概率的映射函数π ( a ∣ s ) \pi(a|s)π(a∣s)。策略不一定非得是神经网络历史上还有表格策略直接把每个状态对应的最优动作记下来但状态一多就存不下和线性策略用一个线性函数做映射表达能力有限。但当状态维度上升到几万甚至上百万比如 Atari 像素或者大模型的文本序列时只有神经网络这种能拟合任意非线性函数的方案才吃得消。这也是为什么本文接下来实现的策略从头到尾都是一个普通的多层感知机MLP。三、第一部分Actor-Critic 网络PPO 属于 Actor-Critic 家族所以第一步是搭建一个同时输出动作概率和状态打分的网络。classActorCritic(nn.Module):def__init__(self,obs_dim4,act_dim2,hidden64):super().__init__()self.actornn.Sequential(nn.Linear(obs_dim,hidden),nn.ReLU(),nn.Linear(hidden,hidden),nn.ReLU(),nn.Linear(hidden,act_dim),)self.criticnn.Sequential(nn.Linear(obs_dim,hidden),nn.ReLU(),nn.Linear(hidden,hidden),nn.ReLU(),nn.Linear(hidden,1),)这里有几个设计细节值得展开讲Actor 和 Critic 用的是两套完全独立的隐藏层而不是共享一套主干后再分叉成两个输出头。这么做是为了避免两个任务之间的梯度互相干扰——Actor 要学的是怎么做对当前局面最有利Critic 要学的是这个局面客观上值多少分两个目标并不完全一致独立参数能让两边各自学得更纯粹。Actor 只有 4 → 64 → 64 → 2 这么小的规模输入 4 维状态输出 2 个动作的打分logits再经过 softmax 就变成向左/向右的概率分布。Critic 结构相同只是最后一层输出 1 个标量代表从当前状态出发未来预计能拿多少总分。代码里还有一个容易被忽略但很关键的细节——正交初始化nn.init.orthogonal_(module.weight,gainnp.sqrt(2))...# actor 输出层用小 gain → 初始策略接近均匀nn.init.orthogonal_(self.actor[-1].weight,gain0.01)Actor 的输出层被特意用一个很小的 gain0.01初始化目的是让训练刚开始时策略接近均匀分布左右各 50%。这保证了智能体在训练初期有足够的探索空间不会因为初始化的偏差过早地锁死在某个动作上。这也是 SB3 默认行为的一部分保持一致才能让自研版本和 SB3 版本训练曲线可比。动作是怎么从 logits 采样出来的呢defget_action(self,obs,deterministicFalse):logits,valueself.forward(obs)disttorch.distributions.Categorical(logitslogits)ifdeterministic:actionlogits.argmax(dim-1)else:actiondist.sample()log_probdist.log_prob(action)returnaction,log_prob,value训练阶段用的是dist.sample()——按概率随机抽取而不是直接选分数最高的那个argmax。哪怕网络认为向右推的概率高达 90%仍然保留 10% 的概率去尝试向左推。这种随机性正是策略持续探索的来源对应训练日志里的策略熵entropy——熵越高说明探索性越强熵逐渐降低则说明策略越来越自信、越来越确定。只有在最终评估阶段才会切换成deterministicTrue直接选概率最大的动作。四、第二部分Rollout —— 收集训练数据网络搭好之后需要让它和环境交互攒够一批数据才能开始训练。这个采集经验的过程叫Rolloutdefcollect_rollout(model,env,num_steps2048):obs,_env.reset()transitions[]for_inrange(num_steps):obs_tensortorch.FloatTensor(obs)withtorch.no_grad():action,log_prob,valuemodel.get_action(obs_tensor)next_obs,reward,terminated,truncated,_env.step(action.item())transitions.append({obs:obs,action:action.item(),log_prob:log_prob.item(),value:value.item(),reward:float(reward),terminated:terminated,truncated:truncated,next_obs:next_obsiftruncatedandnotterminatedelseNone,})obsnext_obsifterminatedortruncated:obs,_env.reset()returntransitions,last_bootstrap逻辑本身很直白每一步都用当前策略选一个动作、执行、记录下这一步的状态、动作、对数概率、Critic 打分和奖励直到攒够 2048 步一次 Rollout 的长度。但这段代码里藏着一个非常容易踩的工程坑它严格区分了terminated杆子真的倒了回合自然结束和truncated走满 500 步的硬性上限杆子其实还立着。这两者在数值上都表现为回合结束但含义完全不同——如果把truncated也当成terminated处理Critic 会被错误地告知这里的未来价值是 0进而让策略学到一个荒谬的结论“撑满 500 步是件坏事”因为紧跟在满分之后的结束被当成了负面信号。这是强化学习工程实践里一个很典型的陷阱正确的做法是只有terminated时才把后续价值置零truncated时要用 Critic 对next_obs的估计做自举bootstrap后面计算 GAE 时会具体看到这一点是怎么实现的。五、第三部分GAE —— 给每一步动作打分有了原始的交互数据接下来要回答一个核心问题这一步动作到底比平均水平好多少这个好多少就是优势AdvantagePPO 用它来决定该强化还是削弱某个动作。优势估计最朴素的做法是单步 TD 误差——“这一步实际拿到的奖励 下一步的预期价值 - 这一步的预期价值”这种算法方差小但依赖 Critic 的准确性偏差较大。另一个极端是用蒙特卡洛方法把整局的实际回报直接算出来偏差小但方差很大。GAEGeneralized Advantage Estimation用一个参数λ \lambdaλ在两者之间做平滑折中记作δ t r t γ V ( s t 1 ) − V ( s t ) \delta_t r_t \gamma V(s_{t1}) - V(s_t)δt​rt​γV(st1​)−V(st​)defcompute_gae(model,transitions,last_bootstrap,gamma0.99,lam0.95):...forstepinreversed(range(n)):ttransitions[step]ift[terminated]:# 真正结束V(s) 0deltarewards[step]-values[step]gaedeltaelift[truncated]:# 时间截断用 V(next_obs) bootstrap但不传播 GAEdeltarewards[step]gamma*bootstrap_values[step]-values[step]gaedeltaelse:# 正常步deltarewards[step]gamma*next_value-values[step]gaedeltagamma*lam*gae next_valuevalues[step]advantages.insert(0,gae)这段代码正是上一节提到的坑的正确解法terminated时直接把未来价值算作 0truncated时用 Critic 对next_obs的预测值做 bootstrap但同样不把 GAE 往前传播因为这一局确实在这里被截断了接下来是全新的一局只有正常步骤才会用gamma * lam * gae把当前的 TD 误差和后续的优势估计组合起来逐步向前传播。λ \lambdaλ的取值决定了折中的方向λ 0 \lambda0λ0时 GAE 退化为单步 TD 误差低方差高偏差λ 1 \lambda1λ1时退化为蒙特卡洛回报低偏差高方差。代码里取λ 0.95 \lambda0.95λ0.95是工程实践中被反复验证过的经验值。算完优势之后代码还做了一步标准化advantages(advantages-advantages.mean())/(advantages.std()1e-8)把这一批数据的优势值减去均值再除以标准差让它们落在一个稳定的数值范围内。这一步在实践中对训练稳定性影响很大——数值范围不稳定会导致梯度更新的幅度忽大忽小。六、第四部分PPO 裁剪更新有了优势估计终于可以更新网络参数了。但这里遇到一个历史性的难题学习率太小训练太慢学习率太大策略可能一步崩溃。在 PPO 之前2015 年 Schulman 等人提出的 TRPOTrust Region Policy Optimization用 KL 散度约束来限制每次更新的幅度理论优美但实现复杂需要计算自然梯度和共轭梯度。2017 年同一批作者提出了 PPO用一个简单得多的裁剪Clipping技巧替代了复杂的信任域约束defppo_update(model,optimizer,transitions,advantages,returns,clip_eps0.2,epochs10,batch_size64):...for_inrange(epochs):forstartinrange(0,len(transitions),batch_size):...logits,valuesmodel(batch_obs)disttorch.distributions.Categorical(logitslogits)new_log_probsdist.log_prob(batch_actions)ratiotorch.exp(new_log_probs-batch_old_log_probs)surr1ratio*batch_advantages surr2torch.clamp(ratio,1-clip_eps,1clip_eps)*batch_advantages policy_loss-torch.min(surr1,surr2).mean()value_loss((values-batch_returns)**2).mean()entropydist.entropy().mean()losspolicy_loss0.5*value_loss-0.0*entropy optimizer.zero_grad()loss.backward()nn.utils.clip_grad_norm_(model.parameters(),0.5)optimizer.step()这段代码是整个 PPO 算法的灵魂拆开来看重要性采样比率ratioratio exp(new_log_probs - old_log_probs)代表更新后的新策略和采集数据时的旧策略对同一个动作的打分之比。ratio 1 说明策略完全没变大于 1 说明新策略更倾向于选这个动作小于 1 则相反。裁剪目标surr1/surr2surr1是朴素的比率 × 优势如果不加约束一旦某个动作的优势特别大ratio 就可能被推向很极端的值导致策略一步跳变太远。surr2把 ratio 强行夹在[ 1 − ϵ , 1 ϵ ] [1-\epsilon, 1\epsilon][1−ϵ,1ϵ]代码里ϵ 0.2 \epsilon0.2ϵ0.2也就是[ 0.8 , 1.2 ] [0.8, 1.2][0.8,1.2]之内再取surr1和surr2中较小的那个作为最终目标。这个取小的操作保证了当优势为正、ratio 想要超过 1.2 时会被裁剪限制住当优势为负、ratio 想要跌破 0.8 时同样会被限制——策略每次只被允许挪动一小步这就是Proximal邻近这个名字的由来。同一批数据要重复训练多个 epoch代码里是 10 轮每轮内部还会打乱顺序切分成多个 mini-batch每批 64 条分别更新这是为了在数据利用率和训练稳定性之间取得平衡——完全丢弃数据用一次太浪费但重复次数太多又容易过拟合到这一小批数据上。总损失是三项的加权和策略损失policy_loss负号是因为优化器默认做梯度下降而我们要最大化目标、价值损失value_lossCritic 预测值和实际回报的均方误差系数 0.5、以及熵奖励项鼓励探索不过这份代码里系数设为 0.0即没有额外的熵正则化。此外clip_grad_norm_(model.parameters(), 0.5)对梯度的整体范数做了裁剪这是防止训练过程中出现异常大梯度、导致参数瞬间爆炸的常规工程手段。每次更新还顺带统计了两个安全监测指标近似 KL 散度old_log_probs - new_log_probs的均值衡量新旧策略差了多少和裁剪比例有多大比例的样本真的触发了 clip可以理解为安全阀触发率。这两个数字后续会体现在训练曲线上健康的训练中KL 散度应该稳定压在 0.001~0.02 之间超过 0.03 就说明这一步更新幅度过猛有崩溃风险裁剪比例正常在 5%~20% 波动长期超过 30% 同样是危险信号。七、整体流程三步循环 学习率衰减把前面三个组件拼起来model.learn()的真面目就是这样一个循环modelActorCritic()optimizeroptim.Adam(model.parameters(),lr3e-4)foriterationinrange(40):# 第一步收集经验数据2048 步transitions,last_bootstrapcollect_rollout(model,env,2048)# 第二步计算 GAE 优势advantages,returnscompute_gae(model,transitions,last_bootstrap)# 第三步PPO 更新同一批数据训练 10 个 epochmetricsppo_update(model,optimizer,transitions,advantages,returns)40 轮迭代每轮先跑 2048 步收集数据约合几局到十几局 CartPole具体局数取决于策略当前的存活能力再用这批数据算优势、更新参数如此往复。代码里还有一个容易被忽视但很实用的细节——学习率线性衰减frac1.0-iteration/total_iterations lr3e-4*fracforparam_groupinoptimizer.param_groups:param_group[lr]lr学习率从3e-4随着迭代次数线性降到接近 0。这么做的直觉是训练早期策略离最优解还很远需要较大的步长快速逼近训练后期策略已经接近收敛用较小的步长做精细微调避免在最优解附近来回震荡。根据课程文档给出的实测对比这个自研 PyTorch 版本正是凭借这个学习率衰减策略比使用恒定学习率的 SB3 版本更快达到满分——PyTorch 版本约 25K 步就首次摸到 500 分SB3 版本则需要约 80K 步才稳定在满分。八、怎么判断训练是否成功跑起来之后怎么知道训练到底有没有效果根据课程 1.2 节的经验总结判断训练健康与否重点看四类指标回合平均奖励ep_rew_mean最直观的指标应该呈现初期在随机策略水平附近震荡 → 中期快速上升 → 后期趋于稳定的三段式曲线。CartPole 中随机策略的水平大约是 20 分左右满分是 500 分。策略熵entropy应该从高到低缓慢下降和奖励曲线形成剪刀交叉——奖励上升的同时熵下降是健康信号如果熵过早地骤降到 0说明策略过早地锁死在了某个可能并非最优的动作模式上。价值损失value loss应该逐步减小说明 Critic 对状态价值的预测越来越准。但要注意value loss 减小只能说明 Critic 变准了不代表策略变好了策略好坏还是要看奖励曲线。KL 散度与裁剪比例KL 散度长期保持在 0.001~0.02 的低位、裁剪比例保持在 5%~20% 区间说明每次更新都是符合 PPO小步快跑设计初衷的温和调整一旦长期突破警戒线就要警惕训练崩溃的风险。如果你想亲自验证这套实现的效果可以直接跑一遍课程仓库里hands-on-modern-rl的两个脚本做对比python1-ppo_cartpole.py# SB3 版本python2-pytorch_ppo.py# 本文讲解的纯 PyTorch 自研版本swanlabwatchswanlog# 在浏览器里查看两者的训练曲线对比按课程文档给出的实测记录两个版本最终都能在 CartPole-v1 上跑到500.0 /- 0.0的评估成绩说明这份不到 200 行的纯 PyTorch 实现确实完整复现出了 SB3 内部 PPO 算法的核心行为。九、小结回顾整篇文章model.learn()这一行代码背后本质上就是反复执行这样一个三步循环用当前策略去环境里收集一批数据 → 用 GAE 给每一步动作算出一个比预期好多少的优势值 → 用 PPO 的裁剪目标函数小步更新网络参数。策略网络从随机初始化开始那些带来更高回报的动作被逐步强化那些导致杆子倾倒的动作被逐步抑制整个过程完全没有人工指定的决策规则驱动学习的只有奖励信号和梯度下降。几个特别值得记住的工程细节Actor 和 Critic 用独立的网络参数避免两个学习目标互相干扰必须严格区分terminated和truncated否则价值函数会在满分截断处被错误置零导致智能体学到活到满分是坏事的荒谬结论GAE 用λ \lambdaλ参数在 TD 误差和蒙特卡洛回报之间做偏差-方差折中PPO 用比率裁剪代替复杂的信任域约束让每次只小幅调整策略这件事变得工程上极其简单。从下一节开始对应课程第 8 章及以后这套 PPO 骨架会被反复复用——DPO 绕开了显式的奖励模型直接学习偏好GRPO 用组内相对优势替代 Critic 网络来省去额外的价值网络开销但无论怎么变化收集数据、估计优势、小步更新这个核心循环始终没有变。理解了 CartPole 上这个不到 200 行的 PPO 实现就等于拿到了理解后面所有大模型对齐算法的钥匙。
返回列表