diff --git a/agent/async_agent.py b/agent/async_agent.py index fe2d738..3e156e4 100644 --- a/agent/async_agent.py +++ b/agent/async_agent.py @@ -25,10 +25,10 @@ def train(id, config, learning_network, target_network): config.logger.debug('worker %d, episode %d, return %f, avg return %f, episode steps %d, total steps %d' % ( id, episode, rewards[-1], np.mean(rewards[-100:]), steps, config.total_steps.value)) -def evaluate(config, task, learning_network): +def evaluate(config, task, learning_network, extra): test_rewards = [] test_points = [] - worker = config.worker(config, learning_network, None) + worker = config.worker(config, learning_network, extra) # config.logger = Logger('./evaluation_log', gym.logger) while True: steps = config.total_steps.value @@ -69,8 +69,16 @@ class AsyncAgent: target_network.load_state_dict(learning_network.state_dict()) os.environ['OMP_NUM_THREADS'] = '1' - args = [(i, config, learning_network, target_network) for i in range(config.num_workers)] - args.append((config, task, learning_network)) + if config.worker == NStepQLearning or config.worker == OneStepQLearning or config.worker == OneStepSarsa: + extra = target_network + elif config.worker == ContinuousAdvantageActorCritic: + state_normalizer = StaticNormalizer(task.state_dim) + reward_normalizer = StaticNormalizer(1) + extra = [state_normalizer, reward_normalizer] + else: + extra = None + args = [(i, config, learning_network, extra) for i in range(config.num_workers)] + args.append((config, task, learning_network, extra)) procs = [mp.Process(target=train, args=args[i]) for i in range(config.num_workers)] procs.append(mp.Process(target=evaluate, args=args[-1])) for p in procs: p.start() diff --git a/async_worker/continuous_actor_critic.py b/async_worker/continuous_actor_critic.py index 2a1eff5..3d41cd3 100644 --- a/async_worker/continuous_actor_critic.py +++ b/async_worker/continuous_actor_critic.py @@ -7,9 +7,10 @@ import numpy as np import torch from torch.autograd import Variable import torch.nn as nn +from utils import * class ContinuousAdvantageActorCritic: - def __init__(self, config, learning_network, target_network): + def __init__(self, config, learning_network, extra): self.config = config self.actor_opt = config.actor_optimizer_fn(learning_network.actor.parameters()) self.critic_opt = config.critic_optimizer_fn(learning_network.critic.parameters()) @@ -20,15 +21,21 @@ class ContinuousAdvantageActorCritic: self.learning_network = learning_network self.counter = 0 + self.shared_state_normalizer = extra[0] + self.state_normalizer = StaticNormalizer(self.task.state_dim) + self.shared_reward_normalizer = extra[1] + self.reward_normalizer = StaticNormalizer(1) + def episode(self, deterministic=False): config = self.config + self.state_normalizer.offline_stats.load(self.shared_state_normalizer.offline_stats) + self.reward_normalizer.offline_stats.load(self.shared_reward_normalizer.offline_stats) state = self.task.reset() - state = config.state_shift_fn(state) + state = self.state_normalizer(state) steps = 0 total_reward = 0 pending = [] - while not config.stop_signal.value and \ - (not config.max_episode_length or steps < config.max_episode_length): + while not config.stop_signal.value: 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(), @@ -36,7 +43,9 @@ class ContinuousAdvantageActorCritic: False) action = self.config.action_shift_fn(action) next_state, reward, terminal, _ = self.task.step(action) - next_state = config.state_shift_fn(next_state) + terminal = (terminal or (config.max_episode_length and steps >= config.max_episode_length)) + next_state = self.state_normalizer(next_state) + # next_state = config.state_shift_fn(next_state) # if deterministic: # self.config.logger.scalar_summary('reward', reward, self.counter) @@ -49,7 +58,7 @@ class ContinuousAdvantageActorCritic: steps += 1 total_reward += reward - reward = config.reward_shift_fn(reward) + reward = np.asscalar(self.reward_normalizer(np.array([reward]))) if deterministic: if terminal: @@ -106,4 +115,10 @@ class ContinuousAdvantageActorCritic: break state = next_state + self.shared_state_normalizer.offline_stats.merge(self.state_normalizer.online_stats) + self.state_normalizer.online_stats.zero() + + self.shared_reward_normalizer.offline_stats.merge(self.reward_normalizer.online_stats) + self.reward_normalizer.online_stats.zero() + return steps, total_reward \ No newline at end of file diff --git a/main.py b/main.py index 8dd6f00..8426da2 100644 --- a/main.py +++ b/main.py @@ -66,13 +66,12 @@ def a3c_cart_pole(): def a3c_pendulum(): config = Config() config.task_fn = lambda: Pendulum() - # config.reward_shift_fn = lambda reward: reward / 10 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.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.network_fn = lambda: DisjointActorCriticNet( + lambda: GaussianActorNet(task.state_dim, task.action_dim), + lambda: GaussianCriticNet(task.state_dim)) config.policy_fn = lambda: GaussianPolicy() config.worker = ContinuousAdvantageActorCritic config.discount = 0.99 @@ -81,7 +80,6 @@ def a3c_pendulum(): config.update_interval = 5 config.test_interval = 1 config.test_repetitions = 5 - # config.entropy_weight = 0.0001 config.entropy_weight = 0 config.gradient_clip = 40 config.logger = Logger('./log', gym.logger) @@ -91,14 +89,12 @@ def a3c_pendulum(): def a3c_walker(): config = Config() config.task_fn = lambda: BipedalWalker() - shifter = Shifter() - config.state_shift_fn = lambda state: shifter(state) 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.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.network_fn = lambda: DisjointActorCriticNet( + lambda: GaussianActorNet(task.state_dim, task.action_dim), + lambda: GaussianCriticNet(task.state_dim)) config.policy_fn = lambda: GaussianPolicy() config.worker = ContinuousAdvantageActorCritic config.discount = 0.99 @@ -107,9 +103,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 - config.gradient_clip = 30 + config.gradient_clip = 40 config.logger = Logger('./log', gym.logger) agent = AsyncAgent(config) agent.run() @@ -343,7 +338,7 @@ if __name__ == '__main__': # async_cart_pole() # a3c_cart_pole() # a3c_pendulum() - a3c_walker() + # a3c_walker() # ddpg_pendulum() # ddpg_walker() # ppo_pendulum() diff --git a/network/network.py b/network/network.py index 39be784..85108f4 100644 --- a/network/network.py +++ b/network/network.py @@ -51,7 +51,7 @@ class VanillaNet(BasicNet): # Base class for actor critic method class ActorCriticNet(BasicNet): - def predict(self, x, _): + def predict(self, x): phi = self.forward(x, True) pre_prob = self.fc_actor(phi) prob = F.softmax(pre_prob) diff --git a/utils/__init__.py b/utils/__init__.py index cf3229d..63ad9fd 100644 --- a/utils/__init__.py +++ b/utils/__init__.py @@ -1,5 +1,5 @@ from config import * -from shifter import * +from normalizer import * try: from tf_logger import Logger except: diff --git a/utils/normalizer.py b/utils/normalizer.py new file mode 100644 index 0000000..a194a14 --- /dev/null +++ b/utils/normalizer.py @@ -0,0 +1,62 @@ +import torch + +class StaticNormalizer: + def __init__(self, o_size): + self.offline_stats = SharedStats(o_size) + self.online_stats = SharedStats(o_size) + + def __call__(self, o_): + o = torch.FloatTensor(o_) + self.online_stats.feed(o) + std = (self.offline_stats.v + 1e-6) ** .5 + o = (o - self.offline_stats.m) / std + return o.numpy().reshape(o_.shape) + +class SharedStats: + def __init__(self, o_size): + self.m = torch.zeros(o_size) + self.v = torch.zeros(o_size) + self.n = torch.zeros(1) + self.m.share_memory_() + self.v.share_memory_() + self.n.share_memory_() + + def feed(self, o): + n = self.n[0] + new_m = self.m * (n / (n + 1)) + o / (n + 1) + self.v.copy_(self.v * (n / (n + 1)) + (o - self.m) * (o - new_m) / (n + 1)) + self.m.copy_(new_m) + self.n.add_(1) + + def zero(self): + self.m.zero_() + self.v.zero_() + self.n.zero_() + + def load(self, stats): + self.m.copy_(stats.m) + self.v.copy_(stats.v) + self.n.copy_(stats.n) + + def merge(self, B): + A = self + n_A = self.n[0] + n_B = B.n[0] + n = n_A + n_B + delta = B.m - A.m + m = A.m + delta * n_B / n + v = A.v * n_A + B.v * n_B + delta * delta * n_A * n_B / n + v /= n + self.m.copy_(m) + self.v.copy_(v) + self.n.add_(B.n) + + def state_dict(self): + return {'m': self.m.numpy(), + 'v': self.v.numpy(), + 'n': self.n.numpy()} + + def load_state_dict(self, saved): + self.m = torch.FloatTensor(saved['m']) + self.v = torch.FloatTensor(saved['v']) + self.n = torch.FloatTensor(saved['n']) \ No newline at end of file diff --git a/utils/shifter.py b/utils/shifter.py deleted file mode 100644 index f966931..0000000 --- a/utils/shifter.py +++ /dev/null @@ -1,28 +0,0 @@ -# Adapted from https://github.com/kvfrans/parallel-trpo/blob/master/utils.py -class Shifter: - def __init__(self, filter_mean=True): - self.m = 0 - self.v = 0 - self.n = 0. - self.filter_mean = filter_mean - - def state_dict(self): - return {'m': self.m, - 'v': self.v, - 'n': self.n} - - def load_state_dict(self, saved): - self.m = saved['m'] - self.v = saved['v'] - self.n = saved['n'] - - def __call__(self, o): - self.m = self.m * (self.n / (self.n + 1)) + o * 1 / (1 + self.n) - self.v = self.v * (self.n / (self.n + 1)) + (o - self.m) ** 2 * 1 / (1 + self.n) - self.std = (self.v + 1e-6) ** .5 # std - self.n += 1 - if self.filter_mean: - o_ = (o - self.m) / self.std - else: - o_ = o / self.std - return o_