diff --git a/agent/A2C_agent.py b/agent/A2C_agent.py index ff62ff8..ead73dc 100644 --- a/agent/A2C_agent.py +++ b/agent/A2C_agent.py @@ -5,16 +5,17 @@ ####################################################################### import numpy as np -import torch.multiprocessing as mp from network import * from utils import * from component import * +from .BaseAgent import * import pickle import os import time -class A2CAgent: +class A2CAgent(BaseAgent): def __init__(self, config): + BaseAgent.__init__(self) self.config = config self.task = config.task_fn() self.network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -25,13 +26,6 @@ class A2CAgent: self.episode_rewards = np.zeros(config.num_workers) self.last_episode_rewards = np.zeros(config.num_workers) - def close(self): - self.task.close() - - def save(self, file_name): - with open(file_name, 'wb') as f: - torch.save(self.network.state_dict(), f) - def iteration(self): config = self.config rollout = [] diff --git a/agent/BaseAgent.py b/agent/BaseAgent.py new file mode 100644 index 0000000..6c8feae --- /dev/null +++ b/agent/BaseAgent.py @@ -0,0 +1,18 @@ +####################################################################### +# 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 torch + +class BaseAgent: + def __init__(self): + pass + + def close(self): + if hasattr(self.task, 'close'): + self.task.close() + + def save(self, filename): + pass diff --git a/agent/CategoricalDQN_agent.py b/agent/CategoricalDQN_agent.py index 9c05401..af8a081 100644 --- a/agent/CategoricalDQN_agent.py +++ b/agent/CategoricalDQN_agent.py @@ -12,9 +12,11 @@ import time import os import pickle import torch +from .BaseAgent import * -class CategoricalDQNAgent: +class CategoricalDQNAgent(BaseAgent): def __init__(self, config): + BaseAgent.__init__(self) self.config = config self.task = config.task_fn() self.learning_network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -102,10 +104,3 @@ class CategoricalDQNAgent: self.config.logger.debug('episode steps %d, episode time %f, time per step %f' % (steps, episode_time, episode_time / float(steps))) return total_reward, steps - - def save(self, file_name): - with open(file_name, 'wb') as f: - torch.save(self.learning_network.state_dict(), f) - - def close(self): - pass diff --git a/agent/DDPG_agent.py b/agent/DDPG_agent.py index 310a825..36bb2e6 100644 --- a/agent/DDPG_agent.py +++ b/agent/DDPG_agent.py @@ -12,9 +12,11 @@ from component import * import pickle import os import time +from .BaseAgent import * -class DDPGAgent: +class DDPGAgent(BaseAgent): def __init__(self, config): + BaseAgent.__init__(self) self.config = config self.task = config.task_fn() self.worker_network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -35,13 +37,6 @@ class DDPGAgent: target_param.data.copy_(target_param.data * (1.0 - self.config.target_network_mix) + param.data * self.config.target_network_mix) - def save(self, file_name): - with open(file_name, 'wb') as f: - torch.save(self.worker_network.state_dict(), f) - - def close(self): - pass - def episode(self, deterministic=False, video_recorder=None): self.random_process.reset_states() state = self.task.reset() diff --git a/agent/DQN_agent.py b/agent/DQN_agent.py index a2b7d5f..2404ee4 100644 --- a/agent/DQN_agent.py +++ b/agent/DQN_agent.py @@ -12,9 +12,11 @@ import time import os import pickle import torch +from .BaseAgent import * -class DQNAgent: +class DQNAgent(BaseAgent): def __init__(self, config): + BaseAgent.__init__(self) self.config = config self.task = config.task_fn() self.learning_network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -79,10 +81,3 @@ class DQNAgent: self.config.logger.debug('episode steps %d, episode time %f, time per step %f' % (steps, episode_time, episode_time / float(steps))) return total_reward, steps - - def save(self, file_name): - with open(file_name, 'wb') as f: - torch.save(self.learning_network.state_dict(), f) - - def close(self): - pass diff --git a/agent/NStepDQN_agent.py b/agent/NStepDQN_agent.py index 5ad05c1..599f676 100644 --- a/agent/NStepDQN_agent.py +++ b/agent/NStepDQN_agent.py @@ -12,9 +12,11 @@ import time import os import pickle import torch +from .BaseAgent import * -class NStepDQNAgent: +class NStepDQNAgent(BaseAgent): def __init__(self, config): + BaseAgent.__init__(self) self.config = config self.task = config.task_fn() self.learning_network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -28,13 +30,6 @@ class NStepDQNAgent: self.episode_rewards = np.zeros(config.num_workers) self.last_episode_rewards = np.zeros(config.num_workers) - def close(self): - self.task.close() - - def save(self, file_name): - with open(file_name, 'wb') as f: - torch.save(self.learning_network.state_dict(), f) - def iteration(self): config = self.config rollout = [] diff --git a/agent/PPO_agent.py b/agent/PPO_agent.py index 6737f60..f82e4da 100644 --- a/agent/PPO_agent.py +++ b/agent/PPO_agent.py @@ -12,9 +12,11 @@ from component import * import pickle import os import time +from .BaseAgent import * -class PPOAgent: +class PPOAgent(BaseAgent): def __init__(self, config): + BaseAgent.__init__(self) self.config = config self.task = config.task_fn() self.actor = config.actor_network_fn(self.task.state_dim, self.task.action_dim) @@ -28,14 +30,6 @@ class PPOAgent: 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 = [] diff --git a/agent/QuantileRegressionDQN_agent.py b/agent/QuantileRegressionDQN_agent.py index 32117e6..1d13705 100644 --- a/agent/QuantileRegressionDQN_agent.py +++ b/agent/QuantileRegressionDQN_agent.py @@ -12,9 +12,11 @@ import time import os import pickle import torch +from .BaseAgent import * -class QuantileRegressionDQNAgent: +class QuantileRegressionDQNAgent(BaseAgent): def __init__(self, config): + BaseAgent.__init__(self) self.config = config self.task = config.task_fn() self.learning_network = config.network_fn(self.task.state_dim, self.task.action_dim) @@ -93,10 +95,3 @@ class QuantileRegressionDQNAgent: self.config.logger.debug('episode steps %d, episode time %f, time per step %f' % (steps, episode_time, episode_time / float(steps))) return total_reward, steps - - def save(self, file_name): - with open(file_name, 'wb') as f: - torch.save(self.learning_network.state_dict(), f) - - def close(self): - pass diff --git a/agent/__init__.py b/agent/__init__.py index 326c8bb..15f8b8f 100644 --- a/agent/__init__.py +++ b/agent/__init__.py @@ -1,4 +1,3 @@ -from .async_agent import * from .DQN_agent import * from .DDPG_agent import * from .A2C_agent import * diff --git a/agent/async_agent.py b/agent/async_agent.py deleted file mode 100644 index b073ec4..0000000 --- a/agent/async_agent.py +++ /dev/null @@ -1,113 +0,0 @@ -####################################################################### -# 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 * -from async_worker import * -import pickle -import os -import time -import sys - -def train(id, config, learning_network, extra): - np.random.seed() - torch.manual_seed(np.random.randint(sys.maxsize)) - worker = config.worker(config, learning_network, extra) - episode = 0 - rewards = [] - while not config.stop_signal.value: - steps, reward = worker.episode() - rewards.append(reward) - if len(rewards) > 100: rewards.pop(0) - 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)) - episode += 1 - -def evaluate(config, task, learning_network, extra): - np.random.seed() - torch.manual_seed(np.random.randint(sys.maxsize)) - test_rewards = [] - test_points = [] - test_wall_times = [] - initial_time = time.time() - worker = config.worker(config, learning_network, extra) - while True: - steps = config.total_steps.value - if config.test_interval and steps % config.test_interval == 0: - worker.worker_network.load_state_dict(learning_network.state_dict()) - with open('data/%s-%s-model-%s.bin' % ( - config.tag, config.worker.__name__, task.name), 'wb') as f: - pickle.dump(learning_network.state_dict(), f) - rewards = np.zeros(config.test_repetitions) - for i in range(config.test_repetitions): - rewards[i] = worker.episode(deterministic=True)[1] - config.logger.info('total steps: %d, averaged return per episode: %f(%f)' % \ - (steps, np.mean(rewards), np.std(rewards) / np.sqrt(config.test_repetitions))) - test_rewards.append(np.mean(rewards)) - test_points.append(steps) - test_wall_times.append(time.time() - initial_time) - with open('data/%s-%s-statistics-%s.bin' % ( - config.tag, config.worker.__name__, task.name), 'wb') as f: - pickle.dump([test_rewards, test_points, test_wall_times], f) - if np.mean(rewards) >= config.success_threshold or (config.max_steps and steps >= config.max_steps): - config.stop_signal.value = True - break - -class AsyncAgent: - def __init__(self, config): - self.config = config - self.config.steps_lock = mp.Lock() - self.config.network_lock = mp.Lock() - self.config.total_steps = mp.Value('i', 0) - self.config.stop_signal = mp.Value('i', False) - - def run(self): - config = self.config - task = config.task_fn() - learning_network = config.network_fn() - learning_network.share_memory() - - os.environ['OMP_NUM_THREADS'] = '1' - if config.worker == NStepQLearning or config.worker == OneStepQLearning or config.worker == OneStepSarsa: - target_network = config.network_fn() - target_network.share_memory() - target_network.load_state_dict(learning_network.state_dict()) - extra = target_network - elif config.worker == ContinuousAdvantageActorCritic \ - or config.worker == ProximalPolicyOptimization\ - or config.worker == DeterministicPolicyGradient: - state_normalizer = StaticNormalizer(task.state_dim) - reward_normalizer = StaticNormalizer(1) - extra = [state_normalizer, reward_normalizer] - if config.worker == DeterministicPolicyGradient: - extra.append(config.replay_fn()) - 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() - while True: - time.sleep(1) - for i, p in enumerate(procs): - if not p.is_alive() and not config.stop_signal.value: - config.logger.warning('Worker %d exited unexpectedly.' % i) - p.terminate() - if i == config.num_workers: - target = evaluate - else: - target = train - procs[i] = mp.Process(target=target, args=args[i]) - procs[i].start() - self.config.logger.warning('Worker %d restarted.' % i) - break - if config.stop_signal.value: - break - for p in procs: p.join() diff --git a/main.py b/main.py index d1e960b..e025dea 100644 --- a/main.py +++ b/main.py @@ -273,8 +273,8 @@ if __name__ == '__main__': # n_step_dqn_pixel_atari('BreakoutNoFrameskip-v4') # dqn_ram_atari('Breakout-ramNoFrameskip-v4') - # ddpg_continuous() - ppo_continuous() + ddpg_continuous() + # ppo_continuous() # acvp.train('PongNoFrameskip-v4')