动手学强化学习 第 12 章 PPO 算法 训练代码
创始人
2024-11-14 22:04:57
0

基于 Hands-on-RL/第12章-PPO算法.ipynb at main · boyu-ai/Hands-on-RL · GitHub

理论 PPO 算法

修改了警告和报错

运行环境

Debian GNU/Linux 12 Python 3.9.19 torch 2.0.1 gym 0.26.2

运行代码

PPO.py

#!/usr/bin/env python   import gym import torch import torch.nn.functional as F import numpy as np import matplotlib.pyplot as plt import rl_utils   class PolicyNet(torch.nn.Module):     def __init__(self, state_dim, hidden_dim, action_dim):         super(PolicyNet, self).__init__()         self.fc1 = torch.nn.Linear(state_dim, hidden_dim)         self.fc2 = torch.nn.Linear(hidden_dim, action_dim)      def forward(self, x):         x = F.relu(self.fc1(x))         return F.softmax(self.fc2(x), dim=1)   class ValueNet(torch.nn.Module):     def __init__(self, state_dim, hidden_dim):         super(ValueNet, self).__init__()         self.fc1 = torch.nn.Linear(state_dim, hidden_dim)         self.fc2 = torch.nn.Linear(hidden_dim, 1)      def forward(self, x):         x = F.relu(self.fc1(x))         return self.fc2(x)   class PPO:     ''' PPO算法,采用截断方式 '''      def __init__(self, state_dim, hidden_dim, action_dim, actor_lr, critic_lr,                  lmbda, epochs, eps, gamma, device):         self.actor = PolicyNet(state_dim, hidden_dim, action_dim).to(device)         self.critic = ValueNet(state_dim, hidden_dim).to(device)         self.actor_optimizer = torch.optim.Adam(self.actor.parameters(),                                                 lr=actor_lr)         self.critic_optimizer = torch.optim.Adam(self.critic.parameters(),                                                  lr=critic_lr)         self.gamma = gamma         self.lmbda = lmbda         self.epochs = epochs  # 一条序列的数据用来训练轮数         self.eps = eps  # PPO中截断范围的参数         self.device = device      def take_action(self, state):         state = torch.tensor(np.array([state]), dtype=torch.float).to(self.device)         probs = self.actor(state)         action_dist = torch.distributions.Categorical(probs)         action = action_dist.sample()         return action.item()      def update(self, transition_dict):         states = torch.tensor(np.array(transition_dict['states']),                               dtype=torch.float).to(self.device)         actions = torch.tensor(transition_dict['actions']).view(-1, 1).to(             self.device)         rewards = torch.tensor(transition_dict['rewards'],                                dtype=torch.float).view(-1, 1).to(self.device)         next_states = torch.tensor(np.array(transition_dict['next_states']),                                    dtype=torch.float).to(self.device)         dones = torch.tensor(transition_dict['dones'],                              dtype=torch.float).view(-1, 1).to(self.device)         td_target = rewards + self.gamma * self.critic(next_states) * (1 -                                                                        dones)         td_delta = td_target - self.critic(states)         advantage = rl_utils.compute_advantage(self.gamma, self.lmbda,                                                td_delta.cpu()).to(self.device)         old_log_probs = torch.log(self.actor(states).gather(1,                                                             actions)).detach()          for _ in range(self.epochs):             log_probs = torch.log(self.actor(states).gather(1, actions))             ratio = torch.exp(log_probs - old_log_probs)             surr1 = ratio * advantage             surr2 = torch.clamp(ratio, 1 - self.eps,                                 1 + self.eps) * advantage  # 截断             actor_loss = torch.mean(-torch.min(surr1, surr2))  # PPO损失函数             critic_loss = torch.mean(                 F.mse_loss(self.critic(states), td_target.detach()))             self.actor_optimizer.zero_grad()             self.critic_optimizer.zero_grad()             actor_loss.backward()             critic_loss.backward()             self.actor_optimizer.step()             self.critic_optimizer.step()   actor_lr = 1e-3 critic_lr = 1e-2 num_episodes = 500 hidden_dim = 128 gamma = 0.98 lmbda = 0.95 epochs = 10 eps = 0.2 device = torch.device("cuda") if torch.cuda.is_available() else torch.device(     "cpu")  env_name = 'CartPole-v1' env = gym.make(env_name) env.reset(seed=0) torch.manual_seed(0) state_dim = env.observation_space.shape[0] action_dim = env.action_space.n agent = PPO(state_dim, hidden_dim, action_dim, actor_lr, critic_lr, lmbda,             epochs, eps, gamma, device)  return_list = rl_utils.train_on_policy_agent(env, agent, num_episodes)  episodes_list = list(range(len(return_list))) plt.plot(episodes_list, return_list) plt.xlabel('Episodes') plt.ylabel('Returns') plt.title('PPO on {}'.format(env_name)) plt.show()  mv_return = rl_utils.moving_average(return_list, 9) plt.plot(episodes_list, mv_return) plt.xlabel('Episodes') plt.ylabel('Returns') plt.title('PPO on {}'.format(env_name)) plt.show() 

rl_utils.py 参考

动手学强化学习 第 11 章 TRPO 算法 训练代码-CSDN博客

相关内容

热门资讯

攻略辅助!新九天作弊系统(辅助... 攻略辅助!新九天作弊系统(辅助)原来一直总是有辅助攻略(哔哩哔哩)进入游戏-大厅左侧-新手福利-激活...
黑科技辅助挂!超级三加一辅助工... 黑科技辅助挂!超级三加一辅助工具!总是是有开挂辅助教程(今日头条)-哔哩哔哩黑科技辅助挂!超级三加一...
绝活辅助!新蜜瓜大厅破解(辅助... 绝活辅助!新蜜瓜大厅破解(辅助)确实是有辅助软件(哔哩哔哩)暗藏猫腻,小编详细说明新蜜瓜大厅破解破解...
黑科技插件!丽水都莱辅助工具试... 您好,丽水都莱辅助工具试用这款游戏可以开挂的,确实是有挂的,需要了解加去威信【136704302】很...
据权威媒体报道!pokemmo... 据权威媒体报道!pokemmo辅助器手机版下载!确实确实有开挂辅助教程(详细教程)-哔哩哔哩1.po...
攻略辅助!逍遥辅助下载地址(辅... 攻略辅助!逍遥辅助下载地址(辅助)原来是真的有辅助方法(哔哩哔哩)1、下载好逍遥辅助下载地址正确养号...
2026版教程!海盗来了辅助哪... 2026版教程!海盗来了辅助哪个好!都是是有开挂辅助挂(有挂技术)-哔哩哔哩1、海盗来了辅助哪个好免...
攻略辅助!情怀游戏辅助器(辅助... 攻略辅助!情怀游戏辅助器(辅助)一贯真的有辅助插件(哔哩哔哩)1、实时情怀游戏辅助器透视辅助更新:用...
演示辅助!蘑菇云辅助(辅助)果... 演示辅助!蘑菇云辅助(辅助)果然存在有辅助攻略(哔哩哔哩)1、该软件可以轻松地帮助玩家将蘑菇云辅助辅...
今天下午!福建大玩家辅助器!总... 今天下午!福建大玩家辅助器!总是存在有开挂辅助神器(有挂秘籍)-哔哩哔哩1、今天下午!福建大玩家辅助...