Files
2018-11-11 21:32:23 +08:00

104 lines
3.9 KiB
Python

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