ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PPO算法代码实战(三):PyTorch从零实现PPO求解CartPole

PPO算法代码实战(三):PyTorch从零实现PPO求解CartPole PPO代码实战PyTorch实现前言一、环境介绍1.1 CartPole-v1 环境1.2 依赖安装二、PPO整体架构回顾三、完整代码实现3.1 第一步定义Actor-Critic网络3.2 第二步定义PPO Agent3.3 第三步训练主循环四、Clip机制可视化五、GAE与偏差方差权衡六、训练结果与调参建议6.1 预期训练效果6.2 常见问题与调参总结前言前面两篇文章我们详细讲解了PPO算法的原理从策略梯度到TRPO再到PPO的Clip机制。光说不练假把式本篇我们就用PyTorch从零实现一个完整的PPO算法在经典的CartPole平衡杆环境上跑通训练流程。这篇文章适合已经了解PPO基本原理想动手写代码实践的同学。我们会实现一个最精简但完整可用的PPO不依赖任何深度强化学习库只用到PyTorch和Gymnasium。一、环境介绍1.1 CartPole-v1 环境CartPole是强化学习入门最经典的环境之一状态空间4维连续向量小车位置、小车速度、杆子角度、杆子角速度动作空间2个离散动作小车向左推、向右推奖励每坚持一步得1分杆子倒了或小车出界就结束目标尽可能长时间保持杆子竖立满分500分这个环境简单但足够展示PPO的完整流程非常适合入门实战。1.2 依赖安装pipinstalltorch gymnasium numpy matplotlib二、PPO整体架构回顾在写代码之前我们先回顾一下PPO的Actor-Critic架构整个PPO训练流程是一个循环采样 → 计算优势 → 多轮更新 → 更新旧策略 → 再采样。三、完整代码实现3.1 第一步定义Actor-Critic网络importtorchimporttorch.nnasnnimporttorch.optimasoptimfromtorch.distributionsimportCategoricalimportgymnasiumasgymimportnumpyasnpfromcollectionsimportdequeclassActorCritic(nn.Module):def__init__(self,state_dim,action_dim,hidden_dim128):super(ActorCritic,self).__init__()# Actor 策略网络输出动作概率分布self.actornn.Sequential(nn.Linear(state_dim,hidden_dim),nn.ReLU(),nn.Linear(hidden_dim,hidden_dim),nn.ReLU(),nn.Linear(hidden_dim,action_dim),nn.Softmax(dim-1))# Critic 价值网络输出状态价值 V(s)self.criticnn.Sequential(nn.Linear(state_dim,hidden_dim),nn.ReLU(),nn.Linear(hidden_dim,hidden_dim),nn.ReLU(),nn.Linear(hidden_dim,1))defforward(self,state):# 前向传播同时输出动作概率和状态价值action_probsself.actor(state)state_valueself.critic(state)returnaction_probs,state_valuedefget_action(self,state):# 根据状态采样动作statetorch.FloatTensor(state).unsqueeze(0)action_probs,state_valueself.forward(state)distCategorical(action_probs)actiondist.sample()returnaction.item(),dist.log_prob(action),state_value.item()代码讲解Actor网络输出softmax概率对应每个动作的选择概率Critic网络输出一个标量即该状态的价值估计get_action()用于与环境交互时采样动作同时返回log_prob和价值供后续训练使用3.2 第二步定义PPO AgentclassPPOAgent:def__init__(self,state_dim,action_dim,lr3e-4,gamma0.99,eps_clip0.2,K_epochs4,entropy_coef0.01,value_coef0.5,):self.gammagamma self.eps_clipeps_clip self.K_epochsK_epochs self.entropy_coefentropy_coef self.value_coefvalue_coef# 当前网络和旧网络self.policyActorCritic(state_dim,action_dim)self.policy_oldActorCritic(state_dim,action_dim)self.policy_old.load_state_dict(self.policy.state_dict())self.optimizeroptim.Adam(self.policy.parameters(),lrlr)self.MseLossnn.MSELoss()# 存储轨迹数据self.states[]self.actions[]self.log_probs[]self.rewards[]self.dones[]defupdate(self):# 计算回报蒙特卡洛回报rewards[]discounted_reward0forreward,doneinzip(reversed(self.rewards),reversed(self.dones)):ifdone:discounted_reward0discounted_rewardrewardself.gamma*discounted_reward rewards.insert(0,discounted_reward)# 归一化回报rewardstorch.FloatTensor(rewards)rewards(rewards-rewards.mean())/(rewards.std()1e-8)# 转换为tensorold_statestorch.FloatTensor(self.states)old_actionstorch.LongTensor(self.actions)old_log_probstorch.FloatTensor(self.log_probs)# 多轮更新PPO的核心for_inrange(self.K_epochs):# 前向计算action_probs,state_valuesself.policy(old_states)distCategorical(action_probs)# 新的log probnew_log_probsdist.log_prob(old_actions)state_valuesstate_values.squeeze()# 计算优势advantagesrewards-state_values.detach()# 概率比ratiostorch.exp(new_log_probs-old_log_probs.detach())# PPO Clip目标surr1ratios*advantages surr2torch.clamp(ratios,1-self.eps_clip,1self.eps_clip)*advantages actor_loss-torch.min(surr1,surr2).mean()# Critic损失critic_lossself.MseLoss(state_values,rewards)# 熵奖励鼓励探索entropydist.entropy().mean()# 总损失total_lossactor_lossself.value_coef*critic_loss-self.entropy_coef*entropy# 反向传播self.optimizer.zero_grad()total_loss.backward()self.optimizer.step()# 更新旧网络self.policy_old.load_state_dict(self.policy.state_dict())# 清空轨迹self.clear_buffer()defclear_buffer(self):self.states[]self.actions[]self.log_probs[]self.rewards[]self.dones[]核心代码讲解概率比计算ratios torch.exp(new_log_probs - old_log_probs)因为log概率相减再exp就是概率比PPO-Clip核心surr1ratios*advantages surr2torch.clamp(ratios,1-eps_clip,1eps_clip)*advantages actor_loss-torch.min(surr1,surr2).mean()这就是我们上一篇讲的Clip机制的代码实现多轮更新for _ in range(K_epochs)同一批数据重复训练K次这是PPO样本效率高的关键。3.3 第三步训练主循环deftrain():envgym.make(CartPole-v1)state_dimenv.observation_space.shape[0]action_dimenv.action_space.n agentPPOAgent(state_dim,action_dim)max_episodes500max_timesteps500update_timestep2000# 每多少步更新一次timestep0reward_historydeque(maxlen20)forepisodeinrange(max_episodes):state,_env.reset()ep_reward0fortinrange(max_timesteps):timestep1# 用旧策略采样动作action,log_prob,state_valueagent.policy_old.get_action(state)# 与环境交互next_state,reward,done,truncated,_env.step(action)donedoneortruncated# 存入轨迹bufferagent.states.append(state)agent.actions.append(action)agent.log_probs.append(log_prob)agent.rewards.append(reward)agent.dones.append(done)statenext_state ep_rewardreward# 达到更新步数就更新iftimestep%update_timestep0:agent.update()ifdone:breakreward_history.append(ep_reward)avg_rewardnp.mean(reward_history)ifepisode%100:print(fEpisode{episode}, 奖励:{ep_reward:.1f}, 平均奖励(近20轮):{avg_reward:.1f})# 提前结束连续10轮平均奖励超过480就收敛了ifavg_reward480andlen(reward_history)10:print(f训练完成Episode{episode}, 平均奖励:{avg_reward:.1f})breakenv.close()if__name____main__:train()四、Clip机制可视化我们在上一篇原理篇中详细讲了PPO的Clip机制下面这张图直观展示了裁剪的效果从图中可以看到当优势为正时概率比超过1 ϵ 1\epsilon1ϵ后目标函数不再上升梯度为0当优势为负时概率比低于1 − ϵ 1-\epsilon1−ϵ后目标函数不再变化梯度为0这种机制自动限制了策略更新的幅度保证训练稳定五、GAE与偏差方差权衡本教程为了简洁使用的是蒙特卡洛回报作为优势估计的基础。在更复杂的任务中通常会使用GAE广义优势估计来平衡偏差和方差λ 0 \lambda0λ0是单步TD估计低偏差高方差λ 1 \lambda1λ1是蒙特卡洛估计高偏差低方差常用值λ 0.95 \lambda0.95λ0.95在两者之间取得平衡六、训练结果与调参建议6.1 预期训练效果在CartPole-v1环境上PPO通常在100-200个episode左右就能收敛到接近满分500分。6.2 常见问题与调参训练不收敛检查学习率是不是太大试试1e-4增大update_timestep让每批数据更多增大K_epochs到10奖励波动大对优势做归一化代码里已经做了回报归一化增大batch size减小clip epsilon到0.1探索不足很快收敛到局部最优增大entropy_coef到0.02检查网络容量是不是太小总结本篇我们用PyTorch从零实现了一个完整的PPO算法Actor-Critic双网络架构Actor输出动作概率Critic输出状态价值PPO-Clip目标函数通过裁剪概率比来限制策略更新幅度多轮epoch复用数据采样一次训练K次提升样本效率熵奖励探索机制鼓励策略保持多样性避免过早收敛这个精简版PPO虽然代码量不大但包含了PPO算法的全部核心要素。理解了这份代码再去看Stable Baselines3等成熟库的实现就会轻松很多。下一篇预告PPO算法进阶四连续动作空间与MuJoCo机器人控制实战我们将把PPO从离散动作空间扩展到连续动作空间并在MuJoCo环境上进行训练。
返回列表