From 3935dfc4287a52fbd33a66290c35c7438f5f0c90 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Wed, 4 Oct 2017 11:03:02 -0600 Subject: [PATCH] Unifying networks for continuous A3C and PPO --- agent/PPO_agent.py | 3 +- async_worker/continuous_actor_critic.py | 36 +++++++------- main.py | 26 ++++++----- network/continuous_action_network.py | 62 ++++++++++--------------- 4 files changed, 56 insertions(+), 71 deletions(-) diff --git a/agent/PPO_agent.py b/agent/PPO_agent.py index 8564360..c5177f3 100644 --- a/agent/PPO_agent.py +++ b/agent/PPO_agent.py @@ -117,7 +117,7 @@ class PPOWorker: 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) if config.entropy_weight: - policy_loss += config.entropy_weight * self.actor_net.kl_loss(std) + policy_loss += -config.entropy_weight * self.actor_net.entropy(std) v = self.critic_net.predict(states) value_loss = 0.5 * (returns - v).pow(2).mean() @@ -127,7 +127,6 @@ class PPOWorker: nn.utils.clip_grad_norm(self.critic_net.parameters(), config.gradient_clip) self.critic_opt.step() - actor_net_old.load_state_dict(self.actor_net.state_dict()) self.actor_opt.zero_grad() policy_loss.backward() diff --git a/async_worker/continuous_actor_critic.py b/async_worker/continuous_actor_critic.py index ed1120b..2a1eff5 100644 --- a/async_worker/continuous_actor_critic.py +++ b/async_worker/continuous_actor_critic.py @@ -11,9 +11,8 @@ import torch.nn as nn class ContinuousAdvantageActorCritic: def __init__(self, config, learning_network, target_network): self.config = config - # self.optimizer = config.optimizer_fn(learning_network.parameters()) - self.optimizer = config.optimizer_fn(learning_network.actor_params) - self.critic_optimizer = config.critic_optimizer_fn(learning_network.critic_params) + self.actor_opt = config.actor_optimizer_fn(learning_network.actor.parameters()) + self.critic_opt = config.critic_optimizer_fn(learning_network.critic.parameters()) self.worker_network = config.network_fn() self.worker_network.load_state_dict(learning_network.state_dict()) self.task = config.task_fn() @@ -28,10 +27,10 @@ class ContinuousAdvantageActorCritic: steps = 0 total_reward = 0 pending = [] - pi = Variable(torch.FloatTensor([np.pi])) while not config.stop_signal.value and \ (not config.max_episode_length or steps < config.max_episode_length): - mean, std, value = self.worker_network.predict(np.stack([state])) + mean, std, log_std = self.worker_network.actor.predict(np.stack([state])) + value = self.worker_network.critic.predict(np.stack([state])) action = self.policy.sample(mean.data.numpy().flatten(), std.data.numpy().flatten(), False) @@ -58,7 +57,7 @@ class ContinuousAdvantageActorCritic: state = next_state continue - pending.append([mean, std, value, action, reward]) + pending.append([mean, std, log_std, value, action, reward]) with config.steps_lock: config.total_steps.value += 1 @@ -71,40 +70,37 @@ class ContinuousAdvantageActorCritic: R = self.worker_network.critic(np.stack([next_state])).data GAE = torch.FloatTensor([[0]]) for i in reversed(range(len(pending))): - mean, std, value, action, reward = pending[i] + mean, std, log_std, value, action, reward = pending[i] if i == len(pending) - 1: delta = reward + config.discount * R - value.data else: - delta = reward + pending[i + 1][2].data - value.data + delta = reward + pending[i + 1][3].data - value.data GAE = config.discount * config.gae_tau * GAE + delta action = Variable(torch.FloatTensor([action])) - log_prob = -(action - mean).pow(2) / (2 * std.pow(2)) -\ - std.log() - 0.5 * (2 * pi).log().expand_as(std) - actor_loss += -torch.sum(log_prob) * Variable(GAE) - entropy = 0.5 + std.log() + 0.5 * (2 * pi).log().expand_as(std) - actor_loss += -config.entropy_weight * entropy.sum() + log_density = self.worker_network.actor.log_density(action, mean, log_std, std) + actor_loss += -torch.sum(log_density) * Variable(GAE) + if config.entropy_weight: + actor_loss += -config.entropy_weight * self.worker_network.actor.entropy(std) R = reward + config.discount * R critic_loss += 0.5 * (Variable(R) - value).pow(2) pending = [] self.worker_network.zero_grad() - self.optimizer.zero_grad() - self.critic_optimizer.zero_grad() + self.actor_opt.zero_grad() + self.critic_opt.zero_grad() actor_loss.backward() critic_loss.backward() - nn.utils.clip_grad_norm(self.worker_network.actor_params, config.gradient_clip) - nn.utils.clip_grad_norm(self.worker_network.critic_params, config.gradient_clip) + nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip) for param, worker_param in zip( self.learning_network.parameters(), self.worker_network.parameters()): if param.grad is not None: break param._grad = worker_param.grad - self.optimizer.step() - self.critic_optimizer.step() + self.actor_opt.step() + self.critic_opt.step() self.worker_network.load_state_dict(self.learning_network.state_dict()) - self.worker_network.reset(terminal) if terminal: break diff --git a/main.py b/main.py index 86a012d..8dd6f00 100644 --- a/main.py +++ b/main.py @@ -66,12 +66,13 @@ def a3c_cart_pole(): def a3c_pendulum(): config = Config() config.task_fn = lambda: Pendulum() - config.reward_shift_fn = lambda reward: reward / 10 + # config.reward_shift_fn = lambda reward: reward / 10 task = config.task_fn() - config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001) + 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: ContinuousActorCriticNet( - task.state_dim, task.action_dim, 2, F.tanh) + config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim) + config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim) + config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn) config.policy_fn = lambda: GaussianPolicy() config.worker = ContinuousAdvantageActorCritic config.discount = 0.99 @@ -80,7 +81,8 @@ def a3c_pendulum(): config.update_interval = 5 config.test_interval = 1 config.test_repetitions = 5 - config.entropy_weight = 0.0001 + # config.entropy_weight = 0.0001 + config.entropy_weight = 0 config.gradient_clip = 40 config.logger = Logger('./log', gym.logger) agent = AsyncAgent(config) @@ -92,10 +94,11 @@ def a3c_walker(): shifter = Shifter() config.state_shift_fn = lambda state: shifter(state) task = config.task_fn() - config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001) + 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: ContinuousActorCriticNet( - task.state_dim, task.action_dim, 1, F.tanh) + config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim) + config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim) + config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn) config.policy_fn = lambda: GaussianPolicy() config.worker = ContinuousAdvantageActorCritic config.discount = 0.99 @@ -104,7 +107,8 @@ def a3c_walker(): config.update_interval = 20 config.test_interval = 1 config.test_repetitions = 5 - config.entropy_weight = 0.01 + # config.entropy_weight = 0.01 + config.entropy_weight = 0 config.gradient_clip = 30 config.logger = Logger('./log', gym.logger) agent = AsyncAgent(config) @@ -338,8 +342,8 @@ if __name__ == '__main__': # dqn_cart_pole() # async_cart_pole() # a3c_cart_pole() - a3c_pendulum() - # a3c_walker() + # a3c_pendulum() + a3c_walker() # ddpg_pendulum() # ddpg_walker() # ppo_pendulum() diff --git a/network/continuous_action_network.py b/network/continuous_action_network.py index 9309436..3a6e292 100644 --- a/network/continuous_action_network.py +++ b/network/continuous_action_network.py @@ -6,43 +6,6 @@ from network import * -class ContinuousActorCriticNet(nn.Module, BasicNet): - def __init__(self, state_dim, action_dim, action_scale, action_gate): - super(ContinuousActorCriticNet, self).__init__() - actor_hidden = 200 - critic_hidden = 100 - self.fc_actor = nn.Linear(state_dim, actor_hidden) - self.fc_mean = nn.Linear(actor_hidden, action_dim) - self.fc_std = nn.Linear(actor_hidden, action_dim) - self.action_scale = action_scale - self.action_gate = action_gate - self.actor_params = list(self.fc_actor.parameters()) + \ - list(self.fc_mean.parameters()) + \ - list(self.fc_std.parameters()) - - self.fc_critic = nn.Linear(state_dim, critic_hidden) - self.fc_value = nn.Linear(critic_hidden, 1) - self.critic_params = list(self.fc_critic.parameters()) + \ - list(self.fc_value.parameters()) - - BasicNet.__init__(self, None, False) - - def predict(self, x): - x = self.to_torch_variable(x) - value = self.critic(x) - - x = F.relu(self.fc_actor(x)) - mean = self.action_scale * self.action_gate(self.fc_mean(x)) - std = F.softplus(self.fc_std(x) + 1e-5) - - return mean, std, value - - def critic(self, x): - x = self.to_torch_variable(x) - x = F.relu(self.fc_critic(x)) - x = self.fc_value(x) - return x - class DDPGActorNet(nn.Module, BasicNet): def __init__(self, state_dim, @@ -184,7 +147,7 @@ class GaussianActorNet(nn.Module, BasicNet): log_density = -(x - mean).pow(2) / (2 * var) - 0.5 * torch.log(2 * Variable(torch.FloatTensor([np.pi])).expand_as(x)) - log_std return log_density.sum(1) - def kl_loss(self, std): + 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): @@ -205,3 +168,26 @@ class GaussianCriticNet(nn.Module, BasicNet): def predict(self, x): return self.forward(x) + +class DisjointActorCriticNet: + def __init__(self, actor_network_fn, critic_network_fn): + self.actor = actor_network_fn() + self.critic = critic_network_fn() + + def state_dict(self): + return [self.actor.state_dict(), self.critic.state_dict()] + + def load_state_dict(self, state_dicts): + self.actor.load_state_dict(state_dicts[0]) + self.critic.load_state_dict(state_dicts[1]) + + def share_memory(self): + self.actor.share_memory() + self.critic.share_memory() + + def parameters(self): + return list(self.actor.parameters()) + list(self.critic.parameters()) + + def zero_grad(self): + self.actor.zero_grad() + self.critic.zero_grad()