diff --git a/agent/A2C_agent.py b/agent/A2C_agent.py index 6e81451..87e3f80 100644 --- a/agent/A2C_agent.py +++ b/agent/A2C_agent.py @@ -53,6 +53,19 @@ class A2CAgent: return total_rewards, steps def episode(self, deterministic=False): + config = self.config + for _ in range(config.iteration_log_interval): + self.iteration(deterministic) + new_episode_counts = np.sum(self.episode_counts) + new_total_rewards = np.sum(self.total_rewards) + avg_reward = (new_total_rewards - self.prev_total_rewards) / \ + (new_episode_counts - self.prev_episode_counts + 1e-5) + self.prev_total_rewards = new_total_rewards + self.prev_episode_counts = new_episode_counts + return avg_reward, config.rollout_length * config.num_workers * \ + config.iteration_log_interval + + def iteration(self, deterministic=False): if deterministic: return self.evaluate() @@ -100,19 +113,9 @@ class A2CAgent: value_loss = config.value_loss_weight * 0.5 * (Variable(returns) - value).pow(2) self.optimizer.zero_grad() - (policy_loss + value_loss).sum().backward() + (policy_loss + value_loss).mean().backward() nn.utils.clip_grad_norm(self.network.parameters(), config.gradient_clip) self.optimizer.step() steps = config.rollout_length * config.num_workers self.total_steps += steps - new_episode_counts = np.sum(self.episode_counts) - new_total_rewards = np.sum(self.total_rewards) - avg_reward = (new_total_rewards - self.prev_total_rewards) / \ - (new_episode_counts - self.prev_episode_counts + 1e-5) - self.prev_total_rewards = new_total_rewards - self.prev_episode_counts = new_episode_counts - return avg_reward, steps - - - diff --git a/agent/DQN_agent.py b/agent/DQN_agent.py index 830065a..91d14b3 100644 --- a/agent/DQN_agent.py +++ b/agent/DQN_agent.py @@ -29,8 +29,6 @@ class DQNAgent: def episode(self, deterministic=False): episode_start_time = time.time() state = self.task.reset() - self.history_buffer = [state] * self.config.history_length - state = np.vstack(self.history_buffer) total_reward = 0.0 steps = 0 while True: @@ -42,9 +40,6 @@ class DQNAgent: else: action = self.policy.sample(value) next_state, reward, done, _ = self.task.step(action) - self.history_buffer.pop(0) - self.history_buffer.append(next_state) - next_state = np.vstack(self.history_buffer) total_reward += np.sum(reward * self.config.reward_weight) reward = self.config.reward_shift_fn(reward) if not deterministic: diff --git a/component/atari_wrapper.py b/component/atari_wrapper.py index acd8e90..b100eda 100644 --- a/component/atari_wrapper.py +++ b/component/atari_wrapper.py @@ -1,4 +1,4 @@ -# This file is copied/apdated from +# This file is apdated from # https://raw.githubusercontent.com/transedward/pytorch-dqn/master/utils/atari_wrapper.py import numpy as np @@ -189,3 +189,20 @@ class NormalizeFrame(gym.Wrapper): def _reset(self): return self._normalize(self.env.reset()) + +class StackFrame(gym.Wrapper): + def __init__(self, env=None, history_length=1): + super(StackFrame, self).__init__(env) + self.history_length = history_length + self.buffer = None + + def _reset(self): + state = self.env.reset() + self.buffer = [state] * self.history_length + return np.vstack(self.buffer) + + def _step(self, action): + state, reward, done, info = self.env.step(action) + self.buffer.pop(0) + self.buffer.append(state) + return np.vstack(self.buffer), reward, done, info diff --git a/component/task.py b/component/task.py index 55f97b1..577439c 100644 --- a/component/task.py +++ b/component/task.py @@ -58,7 +58,7 @@ class LunarLander(BasicTask): class PixelAtari(BasicTask): def __init__(self, name, no_op, frame_skip, normalized_state=True, - frame_size=84, max_steps=10000): + frame_size=84, max_steps=10000, history_length=1): BasicTask.__init__(self, max_steps) self.normalized_state = normalized_state self.name = name @@ -69,7 +69,8 @@ class PixelAtari(BasicTask): env = MaxAndSkipEnv(env, skip=frame_skip) if 'FIRE' in env.unwrapped.get_action_meanings(): env = FireResetEnv(env) - self.env = ProcessFrame(env, frame_size) + env = ProcessFrame(env, frame_size) + self.env = StackFrame(env, history_length) self.action_dim = self.env.action_space.n def normalize_state(self, state): diff --git a/main.py b/main.py index 4c89fec..d3fcc65 100644 --- a/main.py +++ b/main.py @@ -95,7 +95,8 @@ def a2c_cart_pole(): def dqn_pixel_atari(name): config = Config() config.history_length = 4 - config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False) + config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False, + history_length=config.history_length) action_dim = config.task_fn().action_dim config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) config.network_fn = lambda: NatureConvNet(config.history_length, action_dim) @@ -162,23 +163,25 @@ def a3c_pixel_atari(name): def a2c_pixel_atari(name): config = Config() - config.history_length = 1 - config.num_workers = 8 - task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42, max_steps=10000) + config.history_length = 4 + config.num_workers = 5 + task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42, + history_length=config.history_length) config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers) task = config.task_fn() - config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001) + config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.0007, eps=1e-5, alpha=0.99) + # config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001) config.network_fn = lambda: OpenAIActorCriticConvNet( config.history_length, task.task.env.action_space.n, LSTM=False, gpu=True) config.reward_shift_fn = lambda r: np.sign(r) config.policy_fn = SamplePolicy config.discount = 0.99 - config.gae_tau = 1.0 + config.gae_tau = 0.97 config.entropy_weight = 0.01 - config.rollout_length = 20 - config.test_interval = 1000 - config.test_repetitions = 10 - config.value_loss_weight = 0.5 + config.rollout_length = 5 + config.test_interval = 0 + config.iteration_log_interval = 100 + config.gradient_clip = 0.5 config.logger = Logger('./log', logger) run_episodes(A2CAgent(config)) diff --git a/utils/config.py b/utils/config.py index 019a197..cb1324f 100644 --- a/utils/config.py +++ b/utils/config.py @@ -53,3 +53,4 @@ class Config: self.render_episode_freq = 0 self.rollout_length = None self.value_loss_weight = 1.0 + self.iteration_log_interval = 30