Create PPO.py

This commit is contained in:
Johnny He
2018-11-11 21:32:23 +08:00
committed by GitHub
parent 121f5af27a
commit 0028cae2fb
+103
View File
@@ -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