Rewrite PPO

This commit is contained in:
Shangtong Zhang
2018-04-05 22:32:34 -06:00
parent 8ad31c79b8
commit 549e4af3b8
9 changed files with 162 additions and 120 deletions
-1
View File
@@ -46,7 +46,6 @@ class A2CAgent:
rewards = config.reward_shift_fn(rewards)
for i, terminal in enumerate(terminals):
if terminals[i]:
next_states[i] = self.task.reset(i)
self.last_episode_rewards[i] = self.episode_rewards[i]
self.episode_rewards[i] = 0
-1
View File
@@ -48,7 +48,6 @@ class NStepDQNAgent:
rewards = config.reward_shift_fn(rewards)
for i, terminal in enumerate(terminals):
if terminals[i]:
next_states[i] = self.task.reset(i)
self.last_episode_rewards[i] = self.episode_rewards[i]
self.episode_rewards[i] = 0
+110
View File
@@ -0,0 +1,110 @@
#######################################################################
# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) #
# Permission given to modify the code as long as you keep this #
# declaration at the top #
#######################################################################
import numpy as np
import torch.multiprocessing as mp
from network import *
from utils import *
from component import *
import pickle
import os
import time
class PPOAgent:
def __init__(self, config):
self.config = config
self.task = config.task_fn()
self.actor = config.actor_network_fn(self.task.state_dim, self.task.action_dim)
self.critic = config.critic_network_fn(self.task.state_dim, self.task.action_dim)
self.actor_opt = config.actor_optimizer_fn(self.actor.parameters())
self.critic_opt = config.critic_optimizer_fn(self.critic.parameters())
self.total_steps = 0
self.episode_rewards = np.zeros(config.num_workers)
self.last_episode_rewards = np.zeros(config.num_workers)
self.state_normalizer = Normalizer(self.task.state_dim)
self.states = self.task.reset()
self.states = self.state_normalizer(self.states)
def close(self):
self.task.close()
def save(self, file_name):
pass
# with open(file_name, 'wb') as f:
# torch.save(self.network.state_dict(), f)
def iteration(self):
config = self.config
rollout = []
states = self.states
for i in range(config.rollout_length):
mean, std, log_std = self.actor.predict(states)
values = self.critic.predict(states)
dist = torch.distributions.Normal(mean, std)
actions = dist.sample()
log_probs = dist.log_prob(actions).detach()
log_probs = torch.sum(log_probs, dim=1, keepdim=True)
next_states, rewards, terminals, _ = self.task.step(actions.data.cpu().numpy())
self.episode_rewards += rewards
rewards = config.reward_shift_fn(rewards)
for i, terminal in enumerate(terminals):
if terminals[i]:
self.last_episode_rewards[i] = self.episode_rewards[i]
self.episode_rewards[i] = 0
next_states = self.state_normalizer(next_states)
rollout.append([states, values, actions, log_probs, rewards, 1 - terminals])
states = next_states
self.states = states
pending_value = self.critic.predict(states)
rollout.append([states, pending_value, None, None, None, None])
processed_rollout = [None] * (len(rollout) - 1)
advantages = self.actor.tensor(np.zeros((config.num_workers, 1)))
returns = pending_value.data
for i in reversed(range(len(rollout) - 1)):
states, value, actions, log_probs, rewards, terminals = rollout[i]
terminals = self.actor.tensor(terminals).unsqueeze(1)
rewards = self.actor.tensor(rewards).unsqueeze(1)
actions = self.actor.variable(actions)
states = self.actor.variable(states)
next_value = rollout[i + 1][1]
returns = rewards + config.discount * terminals * returns
if not config.use_gae:
advantages = returns - value.data
else:
td_error = rewards + config.discount * terminals * next_value.data - value.data
advantages = advantages * config.gae_tau * config.discount * terminals + td_error
processed_rollout[i] = [states, actions, log_probs, returns, advantages]
states, actions, log_probs_old, returns, advantages = map(lambda x: torch.cat(x, dim=0), zip(*processed_rollout))
advantages = (advantages - advantages.mean()) / advantages.std()
advantages = Variable(advantages)
for k in range(config.optimization_epochs):
mean, std, log_std = self.actor.predict(states)
dist = torch.distributions.Normal(mean, std)
log_probs = dist.log_prob(actions)
log_probs = torch.sum(log_probs, dim=1, keepdim=True)
ratio = (log_probs - log_probs_old).exp()
obj = ratio * advantages
obj_clipped = ratio.clamp(1.0 - self.config.ppo_ratio_clip, 1.0 + self.config.ppo_ratio_clip) * advantages
policy_loss = -torch.min(obj, obj_clipped).mean(0)
v = self.critic.predict(states)
value_loss = 0.5 * (Variable(returns) - v).pow(2).mean()
self.actor_opt.zero_grad()
self.critic_opt.zero_grad()
policy_loss.backward()
value_loss.backward()
nn.utils.clip_grad_norm(self.actor.parameters(), config.gradient_clip)
nn.utils.clip_grad_norm(self.critic.parameters(), config.gradient_clip)
self.actor_opt.step()
self.critic_opt.step()
steps = config.rollout_length * config.num_workers
self.total_steps += steps
+1
View File
@@ -5,3 +5,4 @@ from .A2C_agent import *
from .CategoricalDQN_agent import *
from .NStepDQN_agent import *
from .QuantileRegressionDQN_agent import *
from .PPO_agent import *
+4 -1
View File
@@ -151,7 +151,10 @@ def sub_task(parent_pipe, pipe, task_fn, rank, log_dir):
while True:
op, data = pipe.recv()
if op == 'step':
pipe.send(task.step(data))
ob, reward, done, info = task.step(data)
if done:
ob = task.reset()
pipe.send([ob, reward, done, info])
elif op == 'reset':
pipe.send(task.reset())
elif op == 'exit':
+24 -98
View File
@@ -203,105 +203,32 @@ def dqn_ram_atari(name):
# config.double_q = False
run_episodes(DQNAgent(config))
# def a3c_continuous():
# config = Config()
# config.task_fn = lambda: Pendulum()
# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
# task = config.task_fn()
# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001)
# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
# config.network_fn = lambda: DisjointActorCriticNet(
# # lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=False, action_gate=F.tanh, action_scale=2.0),
# lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=True),
# lambda: GaussianCriticNet(task.state_dim))
# config.policy_fn = lambda: GaussianPolicy()
# config.worker = ContinuousAdvantageActorCritic
# config.discount = 0.99
# config.num_workers = 8
# config.update_interval = 20
# config.test_interval = 1
# config.test_repetitions = 1
# config.entropy_weight = 0
# config.gradient_clip = 40
# config.logger = Logger('./log', logger)
# agent = AsyncAgent(config)
# agent.run()
#
# def p3o_continuous():
# config = Config()
# config.task_fn = lambda: Pendulum()
# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
# # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1')
# task = config.task_fn()
# config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim,
# gpu=-1, unit_std=True)
# config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=-1)
# config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
#
# config.policy_fn = lambda: GaussianPolicy()
# config.replay_fn = lambda: GeneralReplay(memory_size=2048, batch_size=2048)
# config.worker = ProximalPolicyOptimization
# config.discount = 0.99
# config.gae_tau = 0.97
# config.num_workers = 6
# config.test_interval = 1
# config.test_repetitions = 1
# config.entropy_weight = 0
# config.gradient_clip = 20
# config.rollout_length = 10000
# config.optimize_epochs = 1
# config.ppo_ratio_clip = 0.2
# config.logger = Logger('./log', logger)
# agent = AsyncAgent(config)
# agent.run()
#
# def d3pg_continuous():
# config = Config()
# config.task_fn = lambda: Pendulum()
# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
# # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1')
# task = config.task_fn()
# config.actor_network_fn = lambda: DeterministicActorNet(
# task.state_dim, task.action_dim, F.tanh, 2, non_linear=F.relu, batch_norm=False)
# config.critic_network_fn = lambda: DeterministicCriticNet(
# task.state_dim, task.action_dim, non_linear=F.relu, batch_norm=False)
# config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4)
# config.critic_optimizer_fn =\
# lambda params: torch.optim.Adam(params, lr=1e-4)
# config.replay_fn = lambda: SharedReplay(memory_size=1000000, batch_size=64,
# state_shape=(task.state_dim, ), action_shape=(task.action_dim, ))
# config.discount = 0.99
# config.random_process_fn = \
# lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2,
# n_steps_annealing=100000)
# config.worker = DeterministicPolicyGradient
# config.num_workers = 6
# config.min_memory_size = 50
# config.target_network_mix = 0.001
# config.test_interval = 500
# config.test_repetitions = 1
# config.gradient_clip = 20
# config.logger = Logger('./log', logger)
# agent = AsyncAgent(config)
# agent.run()
## continuous control
def ppo_continuous():
config = Config()
config.num_workers = 5
task_fn = lambda log_dir: Pendulum(log_dir=log_dir)
# task_fn = lambda log_dir: Roboschool('RoboschoolInvertedPendulum-v1', log_dir=log_dir)
# task_fn = lambda log_dir: Roboschool('RoboschoolAnt-v1', log_dir=log_dir)
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, log_dir=get_default_log_dir(ppo_continuous.__name__))
config.actor_network_fn = lambda state_dim, action_dim: GaussianActorNet(state_dim, action_dim)
config.critic_network_fn = lambda state_dim, action_dim: GaussianCriticNet(state_dim)
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
config.discount = 0.99
config.use_gae = True
config.gae_tau = 0.97
config.gradient_clip = 0.5
config.rollout_length = 20
config.optimize_epochs = 4
config.ppo_ratio_clip = 0.2
config.logger = Logger('./log', logger)
run_iterations(PPOAgent(config))
def ddpg_continuous():
config = Config()
config.task_fn = lambda: Pendulum()
# config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
# config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
# config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
# config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
# config.task_fn = lambda: Roboschool('RoboschoolReacher-v1')
# config.task_fn = lambda: Roboschool('RoboschoolHopper-v1')
@@ -346,9 +273,8 @@ if __name__ == '__main__':
# n_step_dqn_pixel_atari('BreakoutNoFrameskip-v4')
# dqn_ram_atari('Breakout-ramNoFrameskip-v4')
ddpg_continuous()
# ddpg_continuous()
ppo_continuous()
# dqn_pixel_atari('BreakoutNoFrameskip-v4')
# dqn_ram_atari('Pong-ramNoFrameskip-v4')
# acvp.train('PongNoFrameskip-v4')
+11 -19
View File
@@ -92,7 +92,7 @@ class GaussianActorNet(nn.Module, BasicNet):
action_scale=1,
action_gate=F.tanh,
gpu=-1,
unit_std=True,
# unit_std=True,
hidden_size=64,
non_linear=F.tanh):
super(GaussianActorNet, self).__init__()
@@ -100,12 +100,8 @@ class GaussianActorNet(nn.Module, BasicNet):
self.fc2 = nn.Linear(hidden_size, hidden_size)
self.action_mean = nn.Linear(hidden_size, action_dim)
if unit_std:
self.action_log_std = nn.Parameter(torch.zeros(1, action_dim))
else:
self.action_std = nn.Linear(hidden_size, action_dim)
self.action_log_std = nn.Parameter(torch.zeros(1, action_dim))
self.unit_std = unit_std
self.action_scale = action_scale
self.action_gate = action_gate
self.non_linear = non_linear
@@ -119,24 +115,20 @@ class GaussianActorNet(nn.Module, BasicNet):
mean = self.action_mean(phi)
if self.action_gate is not None:
mean = self.action_scale * self.action_gate(mean)
if self.unit_std:
log_std = self.action_log_std.expand_as(mean)
std = log_std.exp()
else:
std = F.softplus(self.action_std(phi)) + 1e-5
log_std = std.log()
log_std = self.action_log_std.expand_as(mean)
std = log_std.exp()
return mean, std, log_std
def predict(self, x):
return self.forward(x)
def log_density(self, x, mean, log_std, std):
var = std.pow(2)
log_density = -(x - mean).pow(2) / (2 * var + 1e-5) - 0.5 * torch.log(2 * Variable(torch.FloatTensor([np.pi])).expand_as(x)) - log_std
return log_density.sum(1)
def entropy(self, std):
return 0.5 * (1 + (2 * std.pow(2) * np.pi + 1e-5).log()).sum(1).mean()
# def log_density(self, x, mean, log_std, std):
# var = std.pow(2)
# log_density = -(x - mean).pow(2) / (2 * var + 1e-5) - 0.5 * torch.log(2 * Variable(torch.FloatTensor([np.pi])).expand_as(x)) - log_std
# return log_density.sum(1)
#
# def entropy(self, std):
# return 0.5 * (1 + (2 * std.pow(2) * np.pi + 1e-5).log()).sum(1).mean()
class GaussianCriticNet(nn.Module, BasicNet):
def __init__(self,
+1
View File
@@ -57,3 +57,4 @@ class Config:
self.categorical_n_atoms = 51
self.num_quantiles = 10
self.gaussian_noise_scale = 0.3
self.optimization_epochs = 4
+11
View File
@@ -13,6 +13,17 @@ class Normalizer:
self.n = 1.0
def __call__(self, x):
if np.isscalar(x) or len(x.shape) == 1:
return self.nomalize_single(x)
elif len(x.shape) == 2:
new_x = np.zeros(x.shape)
for i in range(x.shape[0]):
new_x[i] = self.nomalize_single(x[i])
return new_x
else:
assert 'Unsupported Shape'
def nomalize_single(self, x):
is_scalar = np.isscalar(x)
if is_scalar:
x = np.asarray([x])