From c79cdb18eada93e51081134ab659b78a7b0b7a44 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Tue, 25 Jul 2017 12:06:54 -0600 Subject: [PATCH] Refactor async agent --- DDPG_agent.py | 6 +- async_agent.py | 122 +++++++++------------------ config.py | 27 ++++++ empty_logger.py | 21 +++++ logger.py | 7 +- main.py | 159 +++++++++++++++++------------------ network.py | 3 +- worker.py | 215 ++++++++++++++++++++++++++++-------------------- 8 files changed, 296 insertions(+), 264 deletions(-) create mode 100644 config.py create mode 100644 empty_logger.py diff --git a/DDPG_agent.py b/DDPG_agent.py index db97f40..67f0804 100644 --- a/DDPG_agent.py +++ b/DDPG_agent.py @@ -23,6 +23,7 @@ class DDPGAgent: random_process_fn, test_interval, test_repetitions, + noise_decay_steps, tag, logger): self.task = task_fn() @@ -46,6 +47,8 @@ class DDPGAgent: self.test_repetitions = test_repetitions self.total_steps = 0 self.tag = tag + self.epsilon = 1.0 + self.d_epsilon = 1.0 / noise_decay_steps def soft_update(self, target, src): for target_param, param in zip(target.parameters(), src.parameters()): @@ -66,12 +69,13 @@ class DDPGAgent: if self.total_steps < self.exploration_steps: action = self.task.random_action() else: - action += self.random_process.sample() + action += max(self.epsilon, 0) * self.random_process.sample() self.logger.histo_summary('noised action', action, self.total_steps) next_state, reward, done, info = self.task.step(action) if not deterministic: self.replay.feed([state, action, reward, next_state, int(done)]) self.total_steps += 1 + self.epsilon -= self.d_epsilon steps += 1 total_reward += reward state = next_state diff --git a/async_agent.py b/async_agent.py index 9b51d5c..73c0a14 100644 --- a/async_agent.py +++ b/async_agent.py @@ -13,125 +13,81 @@ from network import * from worker import * import pickle import os -import traceback import time class AsyncAgent: - def __init__(self, - task_fn, - network_fn, - optimizer_fn, - policy_fn, - worker_fn, - discount, - step_limit, - target_network_update_freq, - n_workers, - update_interval, - test_interval, - test_repetitions, - history_length, - tag, - logger): - self.network_fn = network_fn - self.learning_network = network_fn() - self.learning_network.share_memory() - self.target_network = network_fn() - self.target_network.share_memory() - self.target_network.load_state_dict(self.learning_network.state_dict()) - self.worker_fn = worker_fn + def __init__(self, config): + self.config = config + learning_network = config.network_fn() + learning_network.share_memory() + target_network = config.network_fn() + target_network.share_memory() + target_network.load_state_dict(learning_network.state_dict()) - self.optimizer_fn = optimizer_fn - self.task_fn = task_fn - self.task = self.task_fn() - self.step_limit = step_limit - self.discount = discount - self.optimizer_fn = optimizer_fn - self.target_network_update_freq = target_network_update_freq - self.policy_fn = policy_fn - self.steps_lock = mp.Lock() - self.network_lock = mp.Lock() - self.total_steps = mp.Value('i', 0) - self.stop_signal = mp.Value('i', False) - self.n_workers = n_workers - self.update_interval = update_interval - self.test_interval = test_interval - self.test_repetitions = test_repetitions - self.logger = logger - self.history_length = history_length - self.tag = tag + self.task = config.task_fn() - def deterministic_episode(self, task, network): - state = task.reset() - total_rewards = 0 - steps = 0 - network.reset(True) - while not self.step_limit or steps < self.step_limit: - action_value = network.predict(np.stack([state])) - if self.worker_fn == AdvantageActorCritic: - action_value = action_value[0] - action = np.argmax(action_value.data.numpy().flatten()) - state, reward, terminal, _ = task.step(action) - steps += 1 - total_rewards += reward - if terminal: - break - return total_rewards + self.config.learning_network = learning_network + self.config.target_network = target_network + 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 train(self, id): - worker = self.worker_fn(self) + worker = self.config.worker(self.config) episode = 0 rewards = [] - while True and not self.stop_signal.value: + while not self.config.stop_signal.value: steps, reward = worker.episode() rewards.append(reward) if len(rewards) > 100: rewards.pop(0) - self.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, self.total_steps.value)) + self.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, self.config.total_steps.value)) def save(self, file_name): with open(file_name, 'wb') as f: - pickle.dump(self.learning_network.state_dict(), f) + pickle.dump(self.config.learning_network.state_dict(), f) def evaluate(self, id): test_rewards = [] test_points = [] - test_network = self.network_fn() + worker = self.config.worker(self.config) while True: - steps = self.total_steps.value - if steps % self.test_interval == 0: - test_network.load_state_dict(self.learning_network.state_dict()) - self.save('data/%s%s-model-%s.bin' % (self.tag, self.worker_fn.__name__, self.task.name)) - rewards = np.zeros(self.test_repetitions) - for i in range(self.test_repetitions): - rewards[i] = self.deterministic_episode(self.task, test_network) - self.logger.info('total steps: %d, averaged return per episode: %f(%f)' %\ - (steps, np.mean(rewards), np.std(rewards) / np.sqrt(self.test_repetitions))) + steps = self.config.total_steps.value + if steps % self.config.test_interval == 0: + worker.worker_network.load_state_dict(self.config.learning_network.state_dict()) + self.save('data/%s-%s-model-%s.bin' % ( + self.config.tag, self.config.worker.__name__, self.task.name)) + rewards = np.zeros(self.config.test_repetitions) + for i in range(self.config.test_repetitions): + rewards[i] = worker.episode(deterministic=True)[1] + self.config.logger.info('total steps: %d, averaged return per episode: %f(%f)' %\ + (steps, np.mean(rewards), np.std(rewards) / np.sqrt(self.config.test_repetitions))) test_rewards.append(np.mean(rewards)) test_points.append(steps) - with open('data/%s%s-statistics-%s.bin' % ( - self.tag, self.worker_fn.__name__, self.task.name + with open('data/%s-%s-statistics-%s.bin' % ( + self.config.tag, self.config.worker.__name__, self.task.name ), 'wb') as f: pickle.dump([test_points, test_rewards], f) if np.mean(rewards) > self.task.success_threshold: - self.stop_signal.value = True + self.config.stop_signal.value = True break def run(self): os.environ['OMP_NUM_THREADS'] = '1' - procs = [mp.Process(target=self.train, args=(i, )) for i in range(self.n_workers)] - procs.append(mp.Process(target=self.evaluate, args=(self.n_workers, ))) + procs = [mp.Process(target=self.train, args=(i, )) for i in range(self.config.num_workers)] + procs.append(mp.Process(target=self.evaluate, args=(self.config.num_workers, ))) for p in procs: p.start() while True: time.sleep(1) for i, p in enumerate(procs): - if not p.is_alive() and not self.stop_signal.value: - self.logger.warning('Worker %d exited unexpectedly.' % i) + if not p.is_alive() and not self.config.stop_signal.value: + self.config.logger.warning('Worker %d exited unexpectedly.' % i) p.terminate() procs[i] = mp.Process(target=self.train, args=(i, )) procs[i].start() - self.logger.warning('Worker %d restarted.' % i) + self.config.logger.warning('Worker %d restarted.' % i) break - if self.stop_signal.value: + if self.config.stop_signal.value: break for p in procs: p.join() diff --git a/config.py b/config.py new file mode 100644 index 0000000..2796f5a --- /dev/null +++ b/config.py @@ -0,0 +1,27 @@ +####################################################################### +# 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 # +####################################################################### + +class Config: + def __init__(self): + self.task_fn = None + self.optimizer_fn = None + self.network_fn = None + self.policy_fn = None + self.replay_fn = None + self.discount = 0.99 + self.target_network_update_freq = 0 + self.max_episode_length = 0 + self.exploration_steps = 0 + self.logger = None + self.history_length = 1 + self.test_interval = 100 + self.test_repetitions = 50 + self.double_q = False + self.tag = 'vanilla' + self.num_workers = 1 + self.worker = None + self.update_interval = 1 + self.gradient_clip = 40 diff --git a/empty_logger.py b/empty_logger.py new file mode 100644 index 0000000..bf1a57d --- /dev/null +++ b/empty_logger.py @@ -0,0 +1,21 @@ +import numpy as np + +class Logger(object): + def __init__(self, log_dir, vanilla_logger, skip=False): + """Create a summary writer logging to log_dir.""" + self.info = vanilla_logger.info + self.debug = vanilla_logger.debug + self.warning = vanilla_logger.warning + self.skip = skip + + def scalar_summary(self, tag, value, step): + if self.skip: + return + + def image_summary(self, tag, images, step): + if self.skip: + return + + def histo_summary(self, tag, values, step, bins=1000): + if self.skip: + return diff --git a/logger.py b/logger.py index 442d5a9..2938de5 100644 --- a/logger.py +++ b/logger.py @@ -11,11 +11,12 @@ except ImportError: class Logger(object): - def __init__(self, log_dir, plain_logger, skip=False): + def __init__(self, log_dir, vanilla_logger, skip=False): """Create a summary writer logging to log_dir.""" self.writer = tf.summary.FileWriter(log_dir) - self.info = plain_logger.info - self.debug = plain_logger.debug + self.info = vanilla_logger.info + self.debug = vanilla_logger.debug + self.warning = vanilla_logger.warning self.skip = skip def scalar_summary(self, tag, value, step): diff --git a/main.py b/main.py index 9d22f54..7258978 100644 --- a/main.py +++ b/main.py @@ -4,6 +4,7 @@ from DDPG_agent import * from logger import * import logging from random_process import * +from config import Config def dqn_cart_pole(): config = dict() @@ -28,46 +29,40 @@ def dqn_cart_pole(): agent.run() def async_cart_pole(): - config = dict() - config['task_fn'] = lambda: CartPole() - config['optimizer_fn'] = lambda params: torch.optim.Adam(params, 0.001) - config['network_fn'] = lambda: FCNet([4, 50, 200, 2]) - config['policy_fn'] = lambda: GreedyPolicy(epsilon=0.5, final_step=5000, min_epsilon=0.1) - config['worker_fn'] = OneStepQLearning - # config['worker_fn'] = NStepQLearning - # config['worker_fn'] = OneStepSarsa - config['discount'] = 0.99 - config['target_network_update_freq'] = 200 - config['step_limit'] = 200 - config['n_workers'] = 16 - config['update_interval'] = 6 - config['test_interval'] = 4000 - config['test_repetitions'] = 50 - config['history_length'] = 1 - config['logger'] = Logger('./log', gym.logger) - config['tag'] = '' - agent = AsyncAgent(**config) + config = Config() + config.task_fn= lambda: CartPole() + config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) + config.network_fn = lambda: FCNet([4, 50, 200, 2]) + config.policy_fn = lambda: GreedyPolicy(epsilon=0.5, final_step=5000, min_epsilon=0.1) + # config.worker = OneStepQLearning + # config.worker = NStepQLearning + config.worker = OneStepSarsa + config.discount = 0.99 + config.target_network_update_freq = 200 + config.max_episode_length = 200 + config.num_workers = 16 + config.update_interval = 6 + config.test_interval = 1 + config.test_repetitions = 50 + config.logger = Logger('./log', gym.logger) + agent = AsyncAgent(config) agent.run() def a3c_cart_pole(): - update_interval = 6 - config = dict() - config['task_fn'] = lambda: CartPole() - config['optimizer_fn'] = lambda params: torch.optim.Adam(params, 0.001) - config['network_fn'] = lambda: ActorCriticFCNet([4, 200, 2]) - config['policy_fn'] = SamplePolicy - config['worker_fn'] = AdvantageActorCritic - config['discount'] = 0.99 - config['target_network_update_freq'] = 200 - config['step_limit'] = 200 - config['n_workers'] = 16 - config['update_interval'] = update_interval - config['history_length'] = 1 - config['test_interval'] = 4000 - config['test_repetitions'] = 50 - config['logger'] = Logger('./log', gym.logger) - config['tag'] = '' - agent = AsyncAgent(**config) + config = Config() + config.task_fn = lambda: CartPole() + config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) + config.network_fn = lambda: ActorCriticFCNet([4, 200, 2]) + config.policy_fn = SamplePolicy + config.worker = AdvantageActorCritic + config.discount = 0.99 + config.max_episode_length = 200 + config.num_workers = 16 + config.update_interval = 6 + config.test_interval = 1 + config.test_repetitions = 50 + config.logger = Logger('./log', gym.logger) + agent = AsyncAgent(config) agent.run() def dqn_pixel_atari(name): @@ -95,55 +90,48 @@ def dqn_pixel_atari(name): agent.run() def async_pixel_atari(name): - config = dict() - history_length = 1 - n_actions = 6 - config['task_fn'] = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42) - config['optimizer_fn'] = lambda params: torch.optim.Adam(params, lr=0.0001) - config['network_fn'] = lambda: OpenAIConvNet(history_length, - n_actions) - config['policy_fn'] = lambda: StochasticGreedyPolicy(epsilons=[0.7, 0.7, 0.7], - final_step=2000000, - min_epsilons=[0.1, 0.01, 0.5], - probs=[0.4, 0.3, 0.3]) - # config['worker_fn'] = OneStepQLearning - # config['worker_fn'] = NStepQLearning - config['worker_fn'] = OneStepSarsa - config['discount'] = 0.99 - config['target_network_update_freq'] = 10000 - config['step_limit'] = 10000 - config['n_workers'] = 16 - config['update_interval'] = 20 - config['test_interval'] = 50000 - config['test_repetitions'] = 1 - config['history_length'] = history_length - config['logger'] = Logger('./log', gym.logger) - config['tag'] = '' - agent = AsyncAgent(**config) + config = Config() + config.history_length = 1 + config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42) + task = config.task_fn() + config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001) + config.network_fn = lambda: OpenAIConvNet( + config.history_length, task.env.action_space.n) + config.policy_fn = lambda: StochasticGreedyPolicy( + epsilons=[0.7, 0.7, 0.7], final_step=2000000, min_epsilons=[0.1, 0.01, 0.5], + probs=[0.4, 0.3, 0.3]) + # config.worker = OneStepSarsa + config.worker = NStepQLearning + # config.worker = OneStepQLearning + config.discount = 0.99 + config.target_network_update_freq = 10000 + config.max_episode_length = 10000 + config.num_workers = 16 + config.update_interval = 20 + config.test_interval = 50000 + config.test_repetitions = 1 + config.logger = Logger('./log', gym.logger) + agent = AsyncAgent(config) agent.run() def a3c_pixel_atari(name): - config = dict() - history_length = 1 - n_actions = 6 - config['task_fn'] = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42) - config['optimizer_fn'] = lambda params: torch.optim.Adam(params, lr=0.0001) - config['network_fn'] = lambda: OpenAIActorCriticConvNet(history_length, - n_actions, - LSTM=False) - config['policy_fn'] = SamplePolicy - config['worker_fn'] = AdvantageActorCritic - config['discount'] = 0.99 - config['target_network_update_freq'] = 0 - config['step_limit'] = 10000 - config['n_workers'] = 16 - config['update_interval'] = 20 - config['test_interval'] = 50000 - config['test_repetitions'] = 1 - config['history_length'] = history_length - config['logger'] = Logger('./log', gym.logger) - config['tag'] = '' - agent = AsyncAgent(**config) + config = Config() + config.history_length = 1 + config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42) + task = config.task_fn() + config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001) + config.network_fn = lambda: OpenAIActorCriticConvNet( + config.history_length, task.env.action_space.n, LSTM=True) + config.policy_fn = SamplePolicy + config.worker = AdvantageActorCritic + config.discount = 0.99 + config.max_episode_length = 10000 + config.num_workers = 16 + config.update_interval = 20 + config.test_interval = 50000 + config.test_repetitions = 1 + config.logger = Logger('./log', gym.logger) + agent = AsyncAgent(config) agent.run() def ddpg_pendulum(): @@ -185,6 +173,7 @@ def ddpg_bipedal_walker(): config['step_limit'] = 1000 config['tau'] = 0.001 config['exploration_steps'] = 100 + config['noise_decay_steps'] = 10000 config['random_process_fn'] = \ lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2) config['test_interval'] = 10 @@ -204,11 +193,11 @@ if __name__ == '__main__': # dqn_pixel_atari('PongNoFrameskip-v3') # async_pixel_atari('PongNoFrameskip-v3') - # a3c_pixel_atari('PongNoFrameskip-v3') + a3c_pixel_atari('PongNoFrameskip-v3') # dqn_pixel_atari('BreakoutNoFrameskip-v3') # async_pixel_atari('BreakoutNoFrameskip-v3') # a3c_pixel_atari('BreakoutNoFrameskip-v3') # ddpg_pendulum() - ddpg_bipedal_walker() \ No newline at end of file + # ddpg_bipedal_walker() \ No newline at end of file diff --git a/network.py b/network.py index 3c3f60c..8cb4372 100644 --- a/network.py +++ b/network.py @@ -298,7 +298,8 @@ class DDPGActorNet(nn.Module, BasicNet): x = self.to_torch_variable(x) x = F.relu(self.layer1(x)) x = F.relu(self.layer2(x)) - x = self.output_gate(self.layer3(x)) + x = self.layer3(x) + # x = self.output_gate(self.layer3(x)) return x def predict(self, x, to_numpy=True): diff --git a/worker.py b/worker.py index d5271a4..e2a926a 100644 --- a/worker.py +++ b/worker.py @@ -9,233 +9,267 @@ from torch.autograd import Variable import torch.nn as nn class AdvantageActorCritic: - def __init__(self, agent): - self.agent = agent - self.optimizer = agent.optimizer_fn(agent.learning_network.parameters()) - self.worker_network = agent.network_fn() - self.worker_network.load_state_dict(agent.learning_network.state_dict()) - self.task = agent.task_fn() - self.policy = agent.policy_fn() + def __init__(self, config): + self.config = config + self.optimizer = config.optimizer_fn(config.learning_network.parameters()) + self.worker_network = config.network_fn() + self.worker_network.load_state_dict(config.learning_network.state_dict()) + self.task = config.task_fn() + self.policy = config.policy_fn() - def episode(self): + def episode(self, deterministic=False): + config = self.config state = self.task.reset() steps = 0 total_reward = 0 pending = [] - while True and not self.agent.stop_signal.value: + while not config.stop_signal.value and \ + (not config.max_episode_length or steps < config.max_episode_length): prob, log_prob, value = self.worker_network.predict(np.stack([state])) - action = self.policy.sample(prob.data.numpy().flatten()) + action = self.policy.sample(prob.data.numpy().flatten(), deterministic) next_state, reward, terminal, _ = self.task.step(action) - pending.append([prob, log_prob, value, action, reward]) steps += 1 - with self.agent.steps_lock: - self.agent.total_steps.value += 1 total_reward += reward - if terminal or len(pending) >= self.agent.update_interval: + if deterministic: + if terminal: + break + state = next_state + continue + + pending.append([prob, log_prob, value, action, reward]) + with config.steps_lock: + config.total_steps.value += 1 + + if terminal or len(pending) >= config.update_interval: loss = 0 if terminal: R = torch.FloatTensor([[0]]) else: R = self.worker_network.critic(np.stack([next_state])).data + GAE = torch.FloatTensor([[0]]) for i in reversed(range(len(pending))): prob, log_prob, value, action, reward = pending[i] - R = reward + self.agent.discount * R + R = reward + config.discount * R advantage = Variable(R) - value + GAE = config.discount * GAE + advantage.data loss += 0.5 * advantage.pow(2) - loss += -log_prob.gather(1, Variable(torch.LongTensor([[action]]))) * Variable(advantage.data) + loss += -log_prob.gather(1, Variable(torch.LongTensor([[action]]))) * Variable(GAE) loss += 0.01 * torch.sum(torch.mul(prob, log_prob)) pending = [] self.worker_network.zero_grad() loss.backward() - nn.utils.clip_grad_norm(self.worker_network.parameters(), 40) + nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip) self.optimizer.zero_grad() for param, worker_param in zip( - self.agent.learning_network.parameters(), self.worker_network.parameters()): + config.learning_network.parameters(), self.worker_network.parameters()): param._grad = worker_param.grad.clone() self.optimizer.step() - self.worker_network.load_state_dict(self.agent.learning_network.state_dict()) + self.worker_network.load_state_dict(config.learning_network.state_dict()) self.worker_network.reset(terminal) if terminal: break - else: - state = next_state + state = next_state return steps, total_reward class NStepQLearning: - def __init__(self, agent): - self.agent = agent - self.optimizer = agent.optimizer_fn(agent.learning_network.parameters()) - self.worker_network = agent.network_fn() - self.worker_network.load_state_dict(agent.learning_network.state_dict()) - self.task = agent.task_fn() - self.policy = agent.policy_fn() + def __init__(self, config): + self.config = config + self.optimizer = config.optimizer_fn(config.learning_network.parameters()) + self.worker_network = config.network_fn() + self.worker_network.load_state_dict(config.learning_network.state_dict()) + self.task = config.task_fn() + self.policy = config.policy_fn() - def episode(self): + def episode(self, deterministic=False): + config = self.config state = self.task.reset() steps = 0 total_reward = 0 pending = [] - while True and not self.agent.stop_signal.value: + while not config.stop_signal.value and \ + (not config.max_episode_length or steps < config.max_episode_length): q = self.worker_network.predict(np.stack([state])) - action = self.policy.sample(q.data.numpy().flatten()) + action = self.policy.sample(q.data.numpy().flatten(), deterministic) next_state, reward, terminal, _ = self.task.step(action) - pending.append([q, action, reward]) steps += 1 - with self.agent.steps_lock: - self.agent.total_steps.value += 1 total_reward += reward - if terminal or len(pending) >= self.agent.update_interval: + if deterministic: + if terminal: + break + state = next_state + continue + + with config.steps_lock: + config.total_steps.value += 1 + pending.append([q, action, reward]) + + if terminal or len(pending) >= config.update_interval: loss = 0 if terminal: R = torch.FloatTensor([[0]]) else: - R, _ = self.agent.target_network.predict( + R, _ = config.target_network.predict( np.stack([next_state])).data.max(1) for i in reversed(range(len(pending))): q, action, reward = pending[i] - R = reward + self.agent.discount * R + R = reward + config.discount * R loss += 0.5 * (Variable(R) - q.gather(1, Variable(torch.LongTensor([[action]])))).pow(2) pending = [] self.worker_network.zero_grad() loss.backward() - nn.utils.clip_grad_norm(self.worker_network.parameters(), 40) + nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip) self.optimizer.zero_grad() for param, worker_param in zip( - self.agent.learning_network.parameters(), self.worker_network.parameters()): + config.learning_network.parameters(), self.worker_network.parameters()): param._grad = worker_param.grad.clone() self.optimizer.step() - self.worker_network.load_state_dict(self.agent.learning_network.state_dict()) + self.worker_network.load_state_dict(config.learning_network.state_dict()) self.worker_network.reset(terminal) if terminal: break - else: - state = next_state + state = next_state - if self.agent.total_steps.value % self.agent.target_network_update_freq == 0: - self.agent.target_network.load_state_dict( - self.agent.learning_network.state_dict()) + if config.total_steps.value % config.target_network_update_freq == 0: + config.target_network.load_state_dict(config.learning_network.state_dict()) return steps, total_reward class OneStepQLearning: - def __init__(self, agent): - self.agent = agent - self.optimizer = agent.optimizer_fn(agent.learning_network.parameters()) - self.worker_network = agent.network_fn() - self.worker_network.load_state_dict(agent.learning_network.state_dict()) - self.task = agent.task_fn() - self.policy = agent.policy_fn() + def __init__(self, config): + self.config = config + self.optimizer = config.optimizer_fn(config.learning_network.parameters()) + self.worker_network = config.network_fn() + self.worker_network.load_state_dict(config.learning_network.state_dict()) + self.task = config.task_fn() + self.policy = config.policy_fn() - def episode(self): + def episode(self, deterministic=False): + config = self.config state = self.task.reset() steps = 0 total_reward = 0 pending = [] - while True and not self.agent.stop_signal.value: + while not config.stop_signal.value and \ + (not config.max_episode_length or steps < config.max_episode_length): q = self.worker_network.predict(np.stack([state])) - action = self.policy.sample(q.data.numpy().flatten()) + action = self.policy.sample(q.data.numpy().flatten(), deterministic) next_state, reward, terminal, _ = self.task.step(action) - pending.append([q, action, reward, next_state]) steps += 1 - with self.agent.steps_lock: - self.agent.total_steps.value += 1 total_reward += reward - if terminal or len(pending) >= self.agent.update_interval: + if deterministic: + if terminal: + break + state = next_state + continue + + with config.steps_lock: + config.total_steps.value += 1 + pending.append([q, action, reward, next_state]) + + if terminal or len(pending) >= config.update_interval: loss = 0 for i in range(len(pending)): q, action, reward, next_state = pending[i] - q_next, _ = self.agent.target_network.predict(np.stack([next_state])).data.max(1) + q_next, _ = config.target_network.predict(np.stack([next_state])).data.max(1) if terminal and i == len(pending) - 1: q_next = torch.FloatTensor([[0]]) - q_next = self.agent.discount * q_next + reward + q_next = config.discount * q_next + reward q = q.gather(1, Variable(torch.LongTensor([[action]]))) loss += 0.5 * (q - Variable(q_next)).pow(2) pending = [] self.worker_network.zero_grad() loss.backward() - nn.utils.clip_grad_norm(self.worker_network.parameters(), 40) + nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip) self.optimizer.zero_grad() for param, worker_param in zip( - self.agent.learning_network.parameters(), self.worker_network.parameters()): + config.learning_network.parameters(), self.worker_network.parameters()): param._grad = worker_param.grad.clone() self.optimizer.step() - self.worker_network.load_state_dict(self.agent.learning_network.state_dict()) + self.worker_network.load_state_dict(config.learning_network.state_dict()) self.worker_network.reset(terminal) if terminal: break - else: - state = next_state + state = next_state - if self.agent.total_steps.value % self.agent.target_network_update_freq == 0: - self.agent.target_network.load_state_dict( - self.agent.learning_network.state_dict()) + if config.total_steps.value % config.target_network_update_freq == 0: + config.target_network.load_state_dict(config.learning_network.state_dict()) return steps, total_reward class OneStepSarsa: - def __init__(self, agent): - self.agent = agent - self.optimizer = agent.optimizer_fn(agent.learning_network.parameters()) - self.worker_network = agent.network_fn() - self.worker_network.load_state_dict(agent.learning_network.state_dict()) - self.task = agent.task_fn() - self.policy = agent.policy_fn() + def __init__(self, config): + self.config = config + self.optimizer = config.optimizer_fn(config.learning_network.parameters()) + self.worker_network = config.network_fn() + self.worker_network.load_state_dict(config.learning_network.state_dict()) + self.task = config.task_fn() + self.policy = config.policy_fn() - def episode(self): + def episode(self, deterministic=False): + config = self.config state = self.task.reset() q = self.worker_network.predict(np.stack([state])) - action = self.policy.sample(q.data.numpy().flatten()) + action = self.policy.sample(q.data.numpy().flatten(), deterministic) steps = 0 total_reward = 0 pending = [] - while True and not self.agent.stop_signal.value: + while not config.stop_signal.value and \ + (not config.max_episode_length or steps < config.max_episode_length): next_state, reward, terminal, _ = self.task.step(action) next_q = self.worker_network.predict(np.stack([next_state])) - next_action = self.policy.sample(next_q.data.numpy().flatten()) + next_action = self.policy.sample(next_q.data.numpy().flatten(), deterministic) pending.append([q, action, reward, next_state, next_action]) steps += 1 - with self.agent.steps_lock: - self.agent.total_steps.value += 1 total_reward += reward - if terminal or len(pending) >= self.agent.update_interval: + if deterministic: + if terminal: + break + state = next_state + action = next_action + continue + + with config.steps_lock: + config.total_steps.value += 1 + + if terminal or len(pending) >= config.update_interval: loss = 0 for i in range(len(pending)): q, action, reward, next_state, next_action = pending[i] - q_next = self.agent.target_network.predict(np.stack([next_state])).data + q_next = config.target_network.predict(np.stack([next_state])).data if terminal and i == len(pending) - 1: q_next = torch.FloatTensor([[0]]) else: q_next = q_next.gather(1, torch.LongTensor([[next_action]])) - q_next = self.agent.discount * q_next + reward + q_next = config.discount * q_next + reward q = q.gather(1, Variable(torch.LongTensor([[action]]))) loss += 0.5 * (q - Variable(q_next)).pow(2) pending = [] self.worker_network.zero_grad() loss.backward() - nn.utils.clip_grad_norm(self.worker_network.parameters(), 40) + nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip) self.optimizer.zero_grad() for param, worker_param in zip( - self.agent.learning_network.parameters(), self.worker_network.parameters()): + config.learning_network.parameters(), self.worker_network.parameters()): param._grad = worker_param.grad.clone() self.optimizer.step() - self.worker_network.load_state_dict(self.agent.learning_network.state_dict()) + self.worker_network.load_state_dict(config.learning_network.state_dict()) self.worker_network.reset(terminal) if terminal: @@ -244,8 +278,7 @@ class OneStepSarsa: q = next_q action = next_action - if self.agent.total_steps.value % self.agent.target_network_update_freq == 0: - self.agent.target_network.load_state_dict( - self.agent.learning_network.state_dict()) + if config.total_steps.value % config.target_network_update_freq == 0: + config.target_network.load_state_dict(config.learning_network.state_dict()) return steps, total_reward