From 567e0876c1082be6a8df2716ba160cd869212926 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Fri, 20 Apr 2018 11:37:19 -0600 Subject: [PATCH] Update evaluation scheme --- agent/A2C_agent.py | 4 ++- agent/BaseAgent.py | 42 ++++++++++++++---------- agent/CategoricalDQN_agent.py | 9 +++-- agent/DDPG_agent.py | 11 +++++-- agent/DQN_agent.py | 4 +-- agent/NStepDQN_agent.py | 6 ++-- agent/PPO_agent.py | 2 +- agent/QuantileRegressionDQN_agent.py | 9 +++-- main.py | 8 ++++- network/base_network.py | 2 +- utils/config.py | 1 + utils/normalizer.py | 49 +++++++++++++++++++++++----- 12 files changed, 107 insertions(+), 40 deletions(-) diff --git a/agent/A2C_agent.py b/agent/A2C_agent.py index 424a147..8b3b12d 100644 --- a/agent/A2C_agent.py +++ b/agent/A2C_agent.py @@ -15,7 +15,7 @@ import time class A2CAgent(BaseAgent): def __init__(self, config): - BaseAgent.__init__(self) + BaseAgent.__init__(self, config) self.config = config self.task = config.task_fn() self.network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -80,5 +80,7 @@ class A2CAgent(BaseAgent): nn.utils.clip_grad_norm(self.network.parameters(), config.gradient_clip) self.optimizer.step() + self.evaluate(config.rollout_length) + steps = config.rollout_length * config.num_workers self.total_steps += steps diff --git a/agent/BaseAgent.py b/agent/BaseAgent.py index 00bf40b..a412839 100644 --- a/agent/BaseAgent.py +++ b/agent/BaseAgent.py @@ -8,8 +8,12 @@ import torch import numpy as np class BaseAgent: - def __init__(self): - self.testing = False + def __init__(self, config): + self.config = config + self.evaluation_env = self.config.evaluation_env + if self.evaluation_env is not None: + self.evaluation_state = self.evaluation_env.reset() + self.evaluation_return = 0 def close(self): if hasattr(self.task, 'close'): @@ -22,20 +26,22 @@ class BaseAgent: state_dict = torch.load(filename, map_location=lambda storage, loc: storage) self.network.load_state_dict(state_dict) - def deterministic_test(self): - if self.testing: + def evaluation_action(self, state): + self.config.state_normalizer.set_read_only() + state = self.config.state_normalizer(np.stack([state])) + action = self.network.predict(state, to_numpy=True) + self.config.state_normalizer.unset_read_only() + return np.argmax(action.flatten()) + + def evaluate(self, steps=1): + config = self.config + if config.evaluation_env is None: return - if not self.config.test_interval: - return - if self.total_steps % self.config.test_interval: - return - if not hasattr(self, 'episode'): - return - rewards = [] - self.testing = True - for _ in range(self.config.test_repetitions): - rewards.append(self.episode(deterministic=True)) - self.testing = False - self.config.logger.info('%d deterministic episodes: %f(%f)' % ( - self.config.test_repetitions, np.mean(rewards), np.std(rewards) / np.sqrt(len(rewards)) - )) + for _ in range(steps): + action = self.evaluation_action(self.evaluation_state) + self.evaluation_state, reward, done, _ = self.evaluation_env.step(action) + self.evaluation_return += reward + if done: + self.evaluation_state = self.evaluation_env.reset() + self.config.logger.info('evaluation episode return: %f' % (self.evaluation_return)) + self.evaluation_return = 0 diff --git a/agent/CategoricalDQN_agent.py b/agent/CategoricalDQN_agent.py index c6ef334..9e40447 100644 --- a/agent/CategoricalDQN_agent.py +++ b/agent/CategoricalDQN_agent.py @@ -16,7 +16,7 @@ from .BaseAgent import * class CategoricalDQNAgent(BaseAgent): def __init__(self, config): - BaseAgent.__init__(self) + BaseAgent.__init__(self, config) self.config = config self.task = config.task_fn() self.network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -33,6 +33,11 @@ class CategoricalDQNAgent(BaseAgent): config.categorical_n_atoms)) self.delta_atom = (config.categorical_v_max - config.categorical_v_min) / float(config.categorical_n_atoms - 1) + def evaluation_action(self, state): + value = self.network.predict(np.stack([self.config.state_normalizer(state)])).squeeze(0).data + value = (value * self.atoms).sum(-1).cpu().numpy().flatten() + return np.argmax(value) + def episode(self, deterministic=False): episode_start_time = time.time() state = self.task.reset() @@ -97,7 +102,7 @@ class CategoricalDQNAgent(BaseAgent): nn.utils.clip_grad_norm(self.network.parameters(), self.config.gradient_clip) self.optimizer.step() - self.deterministic_test() + self.evaluate() if not deterministic and self.total_steps % self.config.target_network_update_freq == 0: self.target_network.load_state_dict(self.network.state_dict()) if not deterministic and self.total_steps > self.config.exploration_steps: diff --git a/agent/DDPG_agent.py b/agent/DDPG_agent.py index 43348ed..c588da0 100644 --- a/agent/DDPG_agent.py +++ b/agent/DDPG_agent.py @@ -16,7 +16,7 @@ from .BaseAgent import * class DDPGAgent(BaseAgent): def __init__(self, config): - BaseAgent.__init__(self) + BaseAgent.__init__(self, config) self.config = config self.task = config.task_fn() self.network = DisjointActorCriticNet(self.task.state_dim, self.task.action_dim, @@ -38,6 +38,13 @@ class DDPGAgent(BaseAgent): target_param.data.copy_(target_param.data * (1.0 - self.config.target_network_mix) + param.data * self.config.target_network_mix) + def evaluation_action(self, state): + self.config.state_normalizer.set_read_only() + state = np.stack([self.config.state_normalizer(state)]) + action = self.actor.predict(state, to_numpy=True).flatten() + self.config.state_normalizer.unset_read_only() + return action + def episode(self, deterministic=False, video_recorder=None): self.random_process.reset_states() state = self.task.reset() @@ -69,7 +76,7 @@ class DDPGAgent(BaseAgent): steps += 1 state = next_state - self.deterministic_test() + self.evaluate() if not deterministic and self.replay.size() >= config.min_memory_size: experiences = self.replay.sample() diff --git a/agent/DQN_agent.py b/agent/DQN_agent.py index e6b6f10..dedffb7 100644 --- a/agent/DQN_agent.py +++ b/agent/DQN_agent.py @@ -16,7 +16,7 @@ from .BaseAgent import * class DQNAgent(BaseAgent): def __init__(self, config): - BaseAgent.__init__(self) + BaseAgent.__init__(self, config) self.config = config self.task = config.task_fn() self.network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -74,7 +74,7 @@ class DQNAgent(BaseAgent): nn.utils.clip_grad_norm(self.network.parameters(), self.config.gradient_clip) self.optimizer.step() - self.deterministic_test() + self.evaluate() if not deterministic and self.total_steps % self.config.target_network_update_freq == 0: self.target_network.load_state_dict(self.network.state_dict()) if not deterministic and self.total_steps > self.config.exploration_steps: diff --git a/agent/NStepDQN_agent.py b/agent/NStepDQN_agent.py index d384174..5358f06 100644 --- a/agent/NStepDQN_agent.py +++ b/agent/NStepDQN_agent.py @@ -16,7 +16,7 @@ from .BaseAgent import * class NStepDQNAgent(BaseAgent): def __init__(self, config): - BaseAgent.__init__(self) + BaseAgent.__init__(self, config) self.config = config self.task = config.task_fn() self.network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -72,4 +72,6 @@ class NStepDQNAgent(BaseAgent): self.optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm(self.network.parameters(), config.gradient_clip) - self.optimizer.step() \ No newline at end of file + self.optimizer.step() + + self.evaluate(config.rollout_length) \ No newline at end of file diff --git a/agent/PPO_agent.py b/agent/PPO_agent.py index 7df4bbb..407a6a1 100644 --- a/agent/PPO_agent.py +++ b/agent/PPO_agent.py @@ -16,7 +16,7 @@ from .BaseAgent import * class PPOAgent(BaseAgent): def __init__(self, config): - BaseAgent.__init__(self) + BaseAgent.__init__(self, config) self.config = config self.task = config.task_fn() self.network = config.network_fn(self.task.state_dim, self.task.action_dim) diff --git a/agent/QuantileRegressionDQN_agent.py b/agent/QuantileRegressionDQN_agent.py index df3024d..0448850 100644 --- a/agent/QuantileRegressionDQN_agent.py +++ b/agent/QuantileRegressionDQN_agent.py @@ -16,7 +16,7 @@ from .BaseAgent import * class QuantileRegressionDQNAgent(BaseAgent): def __init__(self, config): - BaseAgent.__init__(self) + BaseAgent.__init__(self, config) self.config = config self.task = config.task_fn() self.network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -35,6 +35,11 @@ class QuantileRegressionDQNAgent(BaseAgent): cond = (x < 1.0).float().detach() return 0.5 * x.pow(2) * cond + (x.abs() - 0.5) * (1 - cond) + def evaluation_action(self, state): + value = self.network.predict(np.stack([self.config.state_normalizer(state)])).squeeze(0).data + value = (value * self.quantile_weight).sum(-1).cpu().numpy().flatten() + return np.argmax(value) + def episode(self, deterministic=False): episode_start_time = time.time() state = self.task.reset() @@ -87,7 +92,7 @@ class QuantileRegressionDQNAgent(BaseAgent): loss.mean(1).sum().backward() self.optimizer.step() - self.deterministic_test() + self.evaluate() if not deterministic and self.total_steps % self.config.target_network_update_freq == 0: self.target_network.load_state_dict(self.network.state_dict()) if not deterministic and self.total_steps > self.config.exploration_steps: diff --git a/main.py b/main.py index 0d927dd..d19c583 100644 --- a/main.py +++ b/main.py @@ -16,6 +16,7 @@ def dqn_cart_pole(): game = 'CartPole-v0' config = Config() config.task_fn = lambda: ClassicalControl(game, max_steps=200) + config.evaluation_env = config.task_fn() config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) config.network_fn = lambda state_dim, action_dim: FCNet(state_dim, 64, action_dim) # config.network_fn = lambda state_dim, action_dim: DuelingFCNet(state_dim, 64, action_dim) @@ -34,6 +35,7 @@ def a2c_cart_pole(): name = 'CartPole-v0' # name = 'MountainCar-v0' task_fn = lambda log_dir: ClassicalControl(name, max_steps=200, log_dir=log_dir) + config.evaluation_env = task_fn(None) config.num_workers = 5 config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, log_dir=get_default_log_dir(a2c_cart_pole.__name__)) @@ -51,6 +53,7 @@ def categorical_dqn_cart_pole(): game = 'CartPole-v0' config = Config() config.task_fn = lambda: ClassicalControl(game, max_steps=200) + config.evaluation_env = config.task_fn() config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) config.network_fn = lambda state_dim, action_dim: \ CategoricalFCNet(state_dim, action_dim, config.categorical_n_atoms) @@ -68,6 +71,7 @@ def categorical_dqn_cart_pole(): def quantile_regression_dqn_cart_pole(): config = Config() config.task_fn = lambda: ClassicalControl('CartPole-v0', max_steps=200) + config.evaluation_env = config.task_fn() config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) config.network_fn = lambda state_dim, action_dim: \ QuantileFCNet(state_dim, action_dim, config.num_quantiles) @@ -83,6 +87,7 @@ def quantile_regression_dqn_cart_pole(): def n_step_dqn_cart_pole(): config = Config() task_fn = lambda log_dir: ClassicalControl('CartPole-v0', max_steps=200, log_dir=log_dir) + config.evaluation_env = task_fn(None) config.num_workers = 5 config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers) config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) @@ -283,7 +288,7 @@ def ppo_continuous(): actor_optimizer_fn = lambda params: torch.optim.Adam(params, 3e-4, eps=1e-5) critic_optimizer_fn = lambda params: torch.optim.Adam(params, 3e-4, eps=1e-5) config.network_fn = lambda state_dim, action_dim: \ - ContinuousActorCriticWrapper(state_dim, action_dim, actor_network_fn, + GaussianActorCriticWrapper(state_dim, action_dim, actor_network_fn, critic_network_fn, actor_optimizer_fn, critic_optimizer_fn) # config.state_normalizer = RunningStatsNormalizer() @@ -310,6 +315,7 @@ def ddpg_continuous(): # config.task_fn = lambda: Roboschool('RoboschoolWalker2d-v1', log_dir=log_dir) # config.task_fn = lambda: DMControl('cartpole', 'balance', log_dir=log_dir) # config.task_fn = lambda: DMControl('finger', 'spin', log_dir=log_dir) + config.evaluation_env = config.task_fn() config.actor_network_fn = lambda state_dim, action_dim: DeterministicActorNet(state_dim, action_dim) config.critic_network_fn = lambda state_dim, action_dim: DeterministicCriticNet(state_dim, action_dim) config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4) diff --git a/network/base_network.py b/network/base_network.py index 01211ad..9ca83b5 100644 --- a/network/base_network.py +++ b/network/base_network.py @@ -146,7 +146,7 @@ class TwoLayerFCNet(nn.Module): y = self.gate(self.fc2(y)) return y -class ContinuousActorCriticWrapper: +class GaussianActorCriticWrapper: def __init__(self, state_dim, action_dim, actor_fn, critic_fn, actor_opt_fn, critic_opt_fn): self.actor = actor_fn(state_dim, action_dim) self.critic = critic_fn(state_dim) diff --git a/utils/config.py b/utils/config.py index 89fb8fd..665fbc9 100644 --- a/utils/config.py +++ b/utils/config.py @@ -58,3 +58,4 @@ class Config: self.num_mini_batches = 32 self.test_interval = 0 self.test_repetitions = 10 + self.evaluation_env = None diff --git a/utils/normalizer.py b/utils/normalizer.py index f28d68f..cfb65ff 100644 --- a/utils/normalizer.py +++ b/utils/normalizer.py @@ -6,9 +6,28 @@ import torch import numpy as np -class RunningStatsNormalizer: - def __init__(self): +class BaseNormalizer: + def __init__(self, read_only=False): + self.read_only = read_only + + def set_read_only(self): + self.read_only = True + + def unset_read_only(self): + self.read_only = False + + def state_dict(self): + return None + + def load_state_dict(self, _): + return + + +class RunningStatsNormalizer(BaseNormalizer): + def __init__(self, read_only=False): + super(RunningStatsNormalizer, self).__init__(read_only) self.needs_reset = True + self.read_only = read_only def reset(self, x_size): self.m = np.zeros(x_size) @@ -16,8 +35,19 @@ class RunningStatsNormalizer: self.n = 0.0 self.needs_reset = False + def state_dict(self): + return {'m': self.m, 'v': self.v, 'n': self.n} + + def load_state_dict(self, stored): + self.m = stored['m'] + self.v = stored['v'] + self.n = stored['n'] + self.needs_reset = False + def __call__(self, x): if np.isscalar(x) or len(x.shape) == 1: + # if dim of x is 1, it can be interpreted as 1 vector entry or batches of scalar entry, + # fortunately resetting the size to 1 applies to both cases if self.needs_reset: self.reset(1) return self.nomalize_single(x) elif len(x.shape) == 2: @@ -33,10 +63,12 @@ class RunningStatsNormalizer: is_scalar = np.isscalar(x) if is_scalar: x = np.asarray([x]) - new_m = self.m * (self.n / (self.n + 1)) + x / (self.n + 1) - self.v = self.v * (self.n / (self.n + 1)) + (x - self.m) * (x - new_m) / (self.n + 1) - self.m = new_m - self.n += 1 + + if not self.read_only: + new_m = self.m * (self.n / (self.n + 1)) + x / (self.n + 1) + self.v = self.v * (self.n / (self.n + 1)) + (x - self.m) * (x - new_m) / (self.n + 1) + self.m = new_m + self.n += 1 std = (self.v + 1e-6) ** .5 x = (x - self.m) / std @@ -44,8 +76,9 @@ class RunningStatsNormalizer: x = np.asscalar(x) return x -class RescaleNormalizer: +class RescaleNormalizer(BaseNormalizer): def __init__(self, coef=1.0): + super(RescaleNormalizer, self).__init__() self.coef = coef def __call__(self, x): @@ -55,6 +88,6 @@ class ImageNormalizer(RescaleNormalizer): def __init__(self): RescaleNormalizer.__init__(self, 1.0 / 255) -class SignNormalizer: +class SignNormalizer(BaseNormalizer): def __call__(self, x): return np.sign(x) \ No newline at end of file