From 0028cae2fbaa490dfa758ff21f9caee713309d7c Mon Sep 17 00:00:00 2001 From: Johnny He <269401927@qq.com> Date: Sun, 11 Nov 2018 21:32:23 +0800 Subject: [PATCH] Create PPO.py --- AlphaGo/PPO.py | 103 +++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 103 insertions(+) create mode 100644 AlphaGo/PPO.py diff --git a/AlphaGo/PPO.py b/AlphaGo/PPO.py new file mode 100644 index 0000000..b9c74db --- /dev/null +++ b/AlphaGo/PPO.py @@ -0,0 +1,103 @@ +import argparse +import pickle +from collections import namedtuple + +import os +import numpy as np +import matplotlib.pyplot as plt + +import gym +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.optim as optim +from torch.distributions import Normal +from torch.utils.data.sampler import BatchSampler, SubsetRandomSampler +from resnet import ResNet + +class PPO(): + clip_param = 0.2 + max_grad_norm = 0.5 + ppo_epoch = 10 + buffer_capacity = 1000 + batch_size = 8 + + def __init__(self): + super(PPO, self).__init__() + self.resnet = ResNet() + self.buffer = [] + self.counter = 0 + self.training_step = 0 + + self.actor_optimizer = optim.Adam(self.actor_net.parameters(), 1e-3) + self.critic_net_optimizer = optim.Adam(self.critic_net.parameters(), 4e-3) + if not os.path.exists('../param'): + os.makedirs('../param/net_param') + os.makedirs('../param/img') + + def select_action(self, state): + state = torch.from_numpy(state).float().unsqueeze(0) + with torch.no_grad(): + mu, sigma = self.actor_net(state) + dist = Normal(mu, sigma) + action = dist.sample() + action_log_prob = dist.log_prob(action) + action = action.clamp(-2, 2) + return action.item(), action_log_prob.item() + + def get_value(self, state): + state = torch.from_numpy(state) + with torch.no_grad(): + value = self.critic_net(state) + return value.item() + + def save_param(self): + torch.save(self.actor_net.state_dict(), '../param/net_param/actor_net' + str(time.time())[:10], +'.pkl') + torch.save(self.critic_net.state_dict(), '../param/net_param/critic_net' + str(time.time())[:10], +'.pkl') + + def store_transition(self, transition): + self.buffer.append(transition) + self.counter += 1 + return counter % self.buffer_capacity == 0 + + def update(self): + self.training_step += 1 + + state = torch.tensor([t.state for t in self.buffer], dtype=torch.float) + action = torch.tensor([t.action for t in self.buffer], dtype=torch.float).view(-1, 1) + reward = torch.tensor([t.reward for t in self.buffer], dtype=torch.float).view(-1, 1) + next_state = torch.tensor([t.next_state for t in self.buffer], dtype=torch.float) + old_action_log_prob = torch.tensor([t.a_log_prob for t in self.buffer], dtype=torch.float).view(-1, 1) + + reward = (reward - reward.mean()) / (reward.std() + 1e-10) + with torch.no_grad(): + target_v = reward + args.gamma * self.critic_net(next_state) + + advantage = (target_v - self.critic_net(state)).detach() + for _ in range(self.ppo_epoch): # iteration ppo_epoch + for index in BatchSampler(SubsetRandomSampler(range(self.buffer_capacity), self.batch_size, True)): + # epoch iteration, PPO core!!! + mu, sigma = self.actor_net(state[index]) + n = Normal(mu, sigma) + action_log_prob = n.log_prob(action[index]) + ratio = torch.exp(action_log_prob - old_action_log_prob) + + L1 = ratio * advantage[index] + L2 = torch.clamp(ratio, 1 - self.clip_param, 1 + self.clip_param) * advantage[index] + action_loss = -torch.min(L1, L2).mean() # MAX->MIN desent + self.actor_optimizer.zero_grad() + action_loss.backward() + nn.utils.clip_grad_norm_(self.actor_net.parameters(), self.max_grad_norm) + self.actor_optimizer.step() + + value_loss = F.smooth_l1_loss(self.critic_net(state[index]), target_v[index]) + self.critic_net_optimizer.zero_grad() + value_loss.backward() + nn.utils.clip_grad_norm_(self.critic_net.parameters(), self.max_grad_norm) + self.critic_net_optimizer.step() + + del self.buffer[:] + + +if __name__ == '__main__': + pass