diff --git a/agent/A2C_agent.py b/agent/A2C_agent.py index 728accf..1d7509f 100644 --- a/agent/A2C_agent.py +++ b/agent/A2C_agent.py @@ -44,7 +44,7 @@ class A2CAgent: steps = 0 while True: prob, _, _ = self.network.predict(np.stack([state])) - action = self.policy.sample(prob.data.numpy().flatten(), True) + action = self.policy.sample(prob.data.cpu().numpy().flatten(), True) state, reward, done, _ = self.evaluator.step(action) total_rewards += reward steps += 1 @@ -61,7 +61,7 @@ class A2CAgent: states = self.states for i in range(config.rollout_length): prob, log_prob, value = self.network.predict(states) - actions = [self.policy.sample(p, deterministic) for p in prob.data.numpy()] + actions = [self.policy.sample(p, deterministic) for p in prob.data.cpu().numpy()] actions = config.action_shift_fn(actions) next_states, rewards, terminals, _ = self.task.step(actions) self.episode_rewards += rewards diff --git a/agent/DDPG_agent.py b/agent/DDPG_agent.py index a21229b..a7d6590 100644 --- a/agent/DDPG_agent.py +++ b/agent/DDPG_agent.py @@ -40,6 +40,9 @@ class DDPGAgent: 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() @@ -61,7 +64,6 @@ class DDPGAgent: next_state, reward, done, info = self.task.step(action) if video_recorder is not None: video_recorder.capture_frame() - done = (done or (config.max_episode_length and steps >= config.max_episode_length)) next_state = self.state_normalizer(next_state) total_reward += reward reward = self.reward_normalizer(reward) diff --git a/agent/DQN_agent.py b/agent/DQN_agent.py index 80c3a16..830065a 100644 --- a/agent/DQN_agent.py +++ b/agent/DQN_agent.py @@ -41,8 +41,7 @@ class DQNAgent: action = np.random.randint(0, len(value)) else: action = self.policy.sample(value) - next_state, reward, done, info = self.task.step(action) - done = (done or (self.config.max_episode_length and steps > self.config.max_episode_length)) + 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) @@ -60,41 +59,20 @@ class DQNAgent: states, actions, rewards, next_states, terminals = experiences states = self.task.normalize_state(states) next_states = self.task.normalize_state(next_states) - if self.config.hybrid_reward: - q_next = self.target_network.predict(next_states, True) - target = [] - for q_next_ in q_next: - if self.config.target_type == self.config.q_target: - target.append(q_next_.detach().max(1)[0]) - elif self.config.target_type == self.config.expected_sarsa_target: - target.append(q_next_.detach().mean(1)) - target = torch.stack(target, dim=1).detach() - terminals = self.learning_network.to_torch_variable(terminals).unsqueeze(1) - rewards = self.learning_network.to_torch_variable(rewards) - target = self.config.discount * target * (1 - terminals) - target.add_(rewards) - q = self.learning_network.predict(states, True) - q_action = [] - actions = self.learning_network.to_torch_variable(actions, 'int64').unsqueeze(1) - for q_ in q: - q_action.append(q_.gather(1, actions)) - q_action = torch.cat(q_action, dim=1) - loss = self.learning_network.criterion(q_action, target) + q_next = self.target_network.predict(next_states, False).detach() + if self.config.double_q: + _, best_actions = self.learning_network.predict(next_states).detach().max(1) + q_next = q_next.gather(1, best_actions.unsqueeze(1)).squeeze(1) else: - q_next = self.target_network.predict(next_states, False).detach() - if self.config.double_q: - _, best_actions = self.learning_network.predict(next_states).detach().max(1) - q_next = q_next.gather(1, best_actions.unsqueeze(1)).squeeze(1) - else: - q_next, _ = q_next.max(1) - terminals = self.learning_network.to_torch_variable(terminals) - rewards = self.learning_network.to_torch_variable(rewards) - q_next = self.config.discount * q_next * (1 - terminals) - q_next.add_(rewards) - actions = self.learning_network.to_torch_variable(actions, 'int64').unsqueeze(1) - q = self.learning_network.predict(states, False) - q = q.gather(1, actions).squeeze(1) - loss = self.criterion(q, q_next) + q_next, _ = q_next.max(1) + terminals = self.learning_network.to_torch_variable(terminals) + rewards = self.learning_network.to_torch_variable(rewards) + q_next = self.config.discount * q_next * (1 - terminals) + q_next.add_(rewards) + actions = self.learning_network.to_torch_variable(actions, 'int64').unsqueeze(1) + q = self.learning_network.predict(states, False) + q = q.gather(1, actions).squeeze(1) + loss = self.criterion(q, q_next) self.optimizer.zero_grad() loss.backward() self.optimizer.step() @@ -110,3 +88,6 @@ class DQNAgent: 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/async_worker/actor_critic.py b/async_worker/actor_critic.py index e821da6..07f1cb1 100644 --- a/async_worker/actor_critic.py +++ b/async_worker/actor_critic.py @@ -29,7 +29,6 @@ class AdvantageActorCritic: prob, log_prob, value = self.worker_network.predict(np.stack([state])) action = self.policy.sample(prob.data.numpy().flatten(), deterministic) next_state, reward, terminal, _ = self.task.step(action) - terminal = (terminal or (self.config.max_episode_length and steps > self.config.max_episode_length)) steps += 1 total_reward += reward diff --git a/async_worker/continuous_actor_critic.py b/async_worker/continuous_actor_critic.py index 31b8abb..b62bc5e 100644 --- a/async_worker/continuous_actor_critic.py +++ b/async_worker/continuous_actor_critic.py @@ -43,7 +43,6 @@ class ContinuousAdvantageActorCritic: False) action = self.config.action_shift_fn(action) next_state, reward, terminal, _ = self.task.step(action) - terminal = (terminal or (config.max_episode_length and steps >= config.max_episode_length)) next_state = self.state_normalizer(next_state) steps += 1 diff --git a/async_worker/dpg.py b/async_worker/dpg.py index 6137c25..59f58da 100644 --- a/async_worker/dpg.py +++ b/async_worker/dpg.py @@ -58,7 +58,6 @@ class DeterministicPolicyGradient: if not deterministic: action += self.random_process.sample() next_state, reward, done, info = self.task.step(action) - done = (done or (config.max_episode_length and steps >= config.max_episode_length)) next_state = self.state_normalizer(next_state) total_reward += reward reward = self.reward_normalizer(reward) diff --git a/async_worker/n_step_q.py b/async_worker/n_step_q.py index 47c4cfb..a13def5 100644 --- a/async_worker/n_step_q.py +++ b/async_worker/n_step_q.py @@ -30,7 +30,6 @@ class NStepQLearning: q = self.worker_network.predict(np.stack([state])) action = self.policy.sample(q.data.numpy().flatten(), deterministic) next_state, reward, terminal, _ = self.task.step(action) - terminal = (terminal or (config.max_episode_length and steps >= config.max_episode_length)) steps += 1 total_reward += reward diff --git a/async_worker/one_step_q.py b/async_worker/one_step_q.py index a37bc4c..cbf9e14 100644 --- a/async_worker/one_step_q.py +++ b/async_worker/one_step_q.py @@ -30,7 +30,6 @@ class OneStepQLearning: q = self.worker_network.predict(np.stack([state])) action = self.policy.sample(q.data.numpy().flatten(), deterministic) next_state, reward, terminal, _ = self.task.step(action) - terminal = (terminal or (config.max_episode_length and steps >= config.max_episode_length)) steps += 1 total_reward += reward diff --git a/async_worker/one_step_sarsa.py b/async_worker/one_step_sarsa.py index 1434108..7d2a5e8 100644 --- a/async_worker/one_step_sarsa.py +++ b/async_worker/one_step_sarsa.py @@ -30,7 +30,6 @@ class OneStepSarsa: pending = [] while not config.stop_signal.value: next_state, reward, terminal, _ = self.task.step(action) - terminal = (terminal or (config.max_episode_length and steps >= config.max_episode_length)) next_q = self.worker_network.predict(np.stack([next_state])) next_action = self.policy.sample(next_q.data.numpy().flatten(), deterministic) pending.append([q, action, reward, next_state, next_action]) diff --git a/async_worker/ppo.py b/async_worker/ppo.py index 02af70a..f5b0294 100644 --- a/async_worker/ppo.py +++ b/async_worker/ppo.py @@ -72,7 +72,6 @@ class ProximalPolicyOptimization: values.append(value) state, reward, done, _ = self.task.step(action) state = self.state_normalizer(state) - done = (done or (config.max_episode_length and episode_length > config.max_episode_length)) batched_rewards += reward batched_steps += 1 diff --git a/component/task.py b/component/task.py index d08a4ea..6554d74 100644 --- a/component/task.py +++ b/component/task.py @@ -32,7 +32,7 @@ class BasicTask: done = (done or self.steps >= self.max_steps) if self.normalized_state: next_state = self.normalize_state(next_state) - return next_state, np.sign(reward), done, info + return next_state, reward, done, info def random_action(self): return self.env.action_space.sample() @@ -59,17 +59,16 @@ class LunarLander(BasicTask): name = 'LunarLander-v2' success_threshold = 200 - def __init__(self): - BasicTask.__init__(self) + def __init__(self, max_steps=sys.maxsize): + BasicTask.__init__(self, max_steps) self.env = gym.make(self.name) class PixelAtari(BasicTask): def __init__(self, name, no_op, frame_skip, normalized_state=True, - frame_size=84, success_threshold=1000): - BasicTask.__init__(self) + frame_size=84, max_steps=sys.maxsize): + BasicTask.__init__(self, max_steps) self.normalized_state = normalized_state self.name = name - self.success_threshold = success_threshold env = gym.make(name) assert 'NoFrameskip' in env.spec.id env = EpisodicLifeEnv(env) @@ -87,107 +86,49 @@ class ContinuousMountainCar(BasicTask): name = 'MountainCarContinuous-v0' success_threshold = 90 - def __init__(self): - BasicTask.__init__(self) + def __init__(self, max_steps=sys.maxsize): + BasicTask.__init__(self, max_steps) self.env = gym.make(self.name) self.max_episode_steps = self.env._max_episode_steps self.env._max_episode_steps = sys.maxsize self.action_dim = self.env.action_space.shape[0] self.state_dim = self.env.observation_space.shape[0] - def step(self, action): - action = np.clip(action, -1, 1) - next_state, reward, done, info = self.env.step(action) - return next_state, reward, done, info - - class Pendulum(BasicTask): name = 'Pendulum-v0' success_threshold = -10 - def __init__(self): - BasicTask.__init__(self) + def __init__(self, max_steps=sys.maxsize): + BasicTask.__init__(self, max_steps) self.env = gym.make(self.name) - self.max_episode_steps = self.env._max_episode_steps - self.env._max_episode_steps = sys.maxsize self.action_dim = self.env.action_space.shape[0] self.state_dim = self.env.observation_space.shape[0] def step(self, action): - action = np.clip(action, -2, 2) - next_state, reward, done, info = self.env.step(action) - return next_state, reward, done, info + return BasicTask.step(self, np.clip(action, -2, 2)) -class BipedalWalker(BasicTask): - name = 'BipedalWalker-v2' - success_threshold = 300 - - def __init__(self): - BasicTask.__init__(self) - self.env = gym.make(self.name) - self.max_episode_steps = self.env._max_episode_steps - self.env._max_episode_steps = sys.maxsize - self.action_dim = self.env.action_space.shape[0] - self.state_dim = self.env.observation_space.shape[0] - - def step(self, action): - action = np.clip(action, -1, 1) - next_state, reward, done, info = self.env.step(action) - return next_state, reward, done, info - -class BipedalWalkerHardcore(BasicTask): - name = 'BipedalWalkerHardcore-v2' - success_threshold = 300 - - def __init__(self): - BasicTask.__init__(self) - self.env = gym.make(self.name) - self.max_episode_steps = self.env._max_episode_steps - self.env._max_episode_steps = sys.maxsize - self.action_dim = self.env.action_space.shape[0] - self.state_dim = self.env.observation_space.shape[0] - - def step(self, action): - action = np.clip(action, -1, 1) - next_state, reward, done, info = self.env.step(action) - return next_state, reward, done, info - -class ContinuousLunarLander(BasicTask): - name = 'LunarLanderContinuous-v2' - success_threshold = 300 - - def __init__(self): - BasicTask.__init__(self) - self.env = gym.make(self.name) - self.max_episode_steps = self.env._max_episode_steps - self.env._max_episode_steps = sys.maxsize - self.action_dim = self.env.action_space.shape[0] - self.state_dim = self.env.observation_space.shape[0] - - def step(self, action): - action = np.clip(action, -1, 1) - next_state, reward, done, info = self.env.step(action) - return next_state, reward, done, info - -class Roboschool(BasicTask): - def __init__(self, name, success_threshold=sys.maxsize, max_episode_steps=None): - import roboschool - BasicTask.__init__(self) +class Box2DContinuous(BasicTask): + def __init__(self, name, max_steps=sys.maxsize): + BasicTask.__init__(self, max_steps) self.name = name self.env = gym.make(self.name) - self.success_threshold = success_threshold - if max_episode_steps is None: - self.max_episode_steps = self.env._max_episode_steps - else: - self.max_episode_steps = max_episode_steps - self.env._max_episode_steps = sys.maxsize self.action_dim = self.env.action_space.shape[0] self.state_dim = self.env.observation_space.shape[0] def step(self, action): - action = np.clip(action, -1, 1) - next_state, reward, done, info = self.env.step(action) - return next_state, reward, done, info + return BasicTask.step(self, np.clip(action, -1, 1)) + +class Roboschool(BasicTask): + def __init__(self, name, success_threshold=sys.maxsize, max_steps=sys.maxsize): + import roboschool + BasicTask.__init__(self, max_steps) + self.name = name + self.env = gym.make(self.name) + self.action_dim = self.env.action_space.shape[0] + self.state_dim = self.env.observation_space.shape[0] + + def step(self, action): + return BasicTask.step(self, np.clip(action, -1, 1)) def sub_task(parent_pipe, pipe, task_fn): parent_pipe.close() diff --git a/main.py b/main.py index 559f8af..e15ac5b 100644 --- a/main.py +++ b/main.py @@ -20,7 +20,6 @@ def dqn_cart_pole(): config.replay_fn = lambda: Replay(memory_size=10000, batch_size=10) config.discount = 0.99 config.target_network_update_freq = 200 - config.max_episode_length = 200 config.exploration_steps = 1000 config.logger = Logger('./log', logger) config.history_length = 2 @@ -41,7 +40,6 @@ def async_cart_pole(): # 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 @@ -156,54 +154,33 @@ def a3c_pixel_atari(name): agent = AsyncAgent(config) agent.run() -def dqn_fruit(): +def a2c_pixel_atari(name): config = Config() - config.task_fn = lambda: Fruit() - config.optimizer_fn = lambda params: torch.optim.SGD(params, 0.01, momentum=0.9) - config.reward_weight = np.ones(10) / 10 - config.hybrid_reward = False - config.network_fn = lambda: FruitHRFCNet(98, 4, config.reward_weight) - config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1) - config.replay_fn = lambda: Replay(memory_size=10000, batch_size=15) - config.discount = 0.95 - config.target_network_update_freq = 200 - config.max_episode_length = 100 - config.exploration_steps = 200 - config.logger = Logger('./log', logger) config.history_length = 1 - config.test_interval = 0 + config.num_workers = 16 + task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42, max_steps=10000) + 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.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 = 0.97 + config.entropy_weight = 0.01 + config.rollout_length = 20 + config.test_interval = 1000 config.test_repetitions = 10 - config.episode_limit = 5000 - config.double_q = False - run_episodes(DQNAgent(config)) - -def hrdqn_fruit(): - config = Config() - config.task_fn = lambda: Fruit(hybrid_reward=True) - config.hybrid_reward = True - config.reward_weight = np.ones(10) / 10 - config.optimizer_fn = lambda params: torch.optim.SGD(params, 0.01, momentum=0.9) - config.network_fn = lambda optimizer_fn: FruitHRFCNet(98, 4, config.reward_weight) - config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1) - config.replay_fn = lambda: HybridRewardReplay(memory_size=10000, batch_size=15) - config.discount = 0.95 - config.target_network_update_freq = 200 - config.max_episode_length = 100 - config.exploration_steps = 200 config.logger = Logger('./log', logger) - config.history_length = 1 - config.test_interval = 0 - config.test_repetitions = 10 - config.target_type = config.expected_sarsa_target - # config.target_type = config.q_target - config.double_q = False - config.episode_limit = 5000 - run_episodes(DQNAgent(config)) + run_episodes(A2CAgent(config)) def a3c_continuous(): config = Config() config.task_fn = lambda: Pendulum() - # config.task_fn = lambda: BipedalWalkerHardcore() + # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') + # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') + # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') 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) @@ -214,7 +191,6 @@ def a3c_continuous(): config.policy_fn = lambda: GaussianPolicy() config.worker = ContinuousAdvantageActorCritic config.discount = 0.99 - config.max_episode_length = task.max_episode_steps config.num_workers = 8 config.update_interval = 20 config.test_interval = 1 @@ -228,8 +204,9 @@ def a3c_continuous(): def p3o_continuous(): config = Config() config.task_fn = lambda: Pendulum() - # config.task_fn = lambda: BipedalWalker() - # config.task_fn = lambda: BipedalWalkerHardcore() + # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') + # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') + # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1') task = config.task_fn() @@ -248,7 +225,6 @@ def p3o_continuous(): config.num_workers = 6 config.test_interval = 1 config.test_repetitions = 1 - config.max_episode_length = task.max_episode_steps config.entropy_weight = 0 config.gradient_clip = 20 config.rollout_length = 10000 @@ -261,10 +237,11 @@ def p3o_continuous(): def d3pg_continuous(): config = Config() config.task_fn = lambda: Pendulum() - # config.task_fn = lambda: ContinuousLunarLander() + # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') + # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') + # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') - # config.task_fn = lambda: BipedalWalker() task = config.task_fn() config.actor_network_fn = lambda: DeterministicActorNet( task.state_dim, task.action_dim, F.tanh, 2, non_linear=F.relu, batch_norm=False) @@ -277,7 +254,6 @@ def d3pg_continuous(): config.replay_fn = lambda: SharedReplay(memory_size=1000000, batch_size=64, state_shape=(task.state_dim, ), action_shape=(task.action_dim, )) config.discount = 0.99 - config.max_episode_length = task.max_episode_steps config.random_process_fn = \ lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2, n_steps_annealing=100000) @@ -294,14 +270,15 @@ def d3pg_continuous(): def ddpg_continuous(): config = Config() - # config.task_fn = lambda: Pendulum() - # config.task_fn = lambda: ContinuousLunarLander() + config.task_fn = lambda: Pendulum() + # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') + # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') + # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') - config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') + # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') # config.task_fn = lambda: Roboschool('RoboschoolHopper-v1') # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1') # config.task_fn = lambda: Roboschool('RoboschoolWalker2d-v1') - # config.task_fn = lambda: BipedalWalker() task = config.task_fn() config.actor_network_fn = lambda: DeterministicActorNet( task.state_dim, task.action_dim, F.tanh, 1, non_linear=F.relu, batch_norm=False, gpu=False) @@ -313,7 +290,6 @@ def ddpg_continuous(): lambda params: torch.optim.Adam(params, lr=1e-3, weight_decay=0.01) config.replay_fn = lambda: HighDimActionReplay(memory_size=1000000, batch_size=64) config.discount = 0.99 - config.max_episode_length = task.max_episode_steps config.random_process_fn = \ lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2, n_steps_annealing=100000) @@ -335,21 +311,19 @@ if __name__ == '__main__': # logger.setLevel(logging.DEBUG) logger.setLevel(logging.INFO) - # dqn_cart_pole() + dqn_cart_pole() # async_cart_pole() # a3c_cart_pole() - a2c_cart_pole() + # a2c_cart_pole() # a3c_continuous() # p3o_continuous() # d3pg_continuous() # ddpg_continuous() - # dqn_fruit() - # hrdqn_fruit() - # dqn_pixel_atari('PongNoFrameskip-v4') # async_pixel_atari('PongNoFrameskip-v4') # a3c_pixel_atari('PongNoFrameskip-v4') + # a2c_pixel_atari('PongNoFrameskip-v4') # dqn_pixel_atari('BreakoutNoFrameskip-v4') # async_pixel_atari('BreakoutNoFrameskip-v4') diff --git a/network/conv_network.py b/network/conv_network.py index 146efdf..b43da71 100644 --- a/network/conv_network.py +++ b/network/conv_network.py @@ -79,7 +79,8 @@ class OpenAIActorCriticConvNet(nn.Module, ActorCriticNet): def __init__(self, in_channels, n_actions, - LSTM=False): + LSTM=False, + gpu=True): super(OpenAIActorCriticConvNet, self).__init__() self.conv1 = nn.Conv2d(in_channels, 32, 3, stride=2, padding=1) self.conv2 = nn.Conv2d(32, 32, 3, stride=2, padding=1) @@ -96,7 +97,7 @@ class OpenAIActorCriticConvNet(nn.Module, ActorCriticNet): self.fc_actor = nn.Linear(hidden_units, n_actions) self.fc_critic = nn.Linear(hidden_units, 1) - BasicNet.__init__(self, gpu=False, LSTM=LSTM) + BasicNet.__init__(self, gpu=gpu, LSTM=LSTM) if LSTM: self.h = self.to_torch_variable(np.zeros((1, hidden_units))) self.c = self.to_torch_variable(np.zeros((1, hidden_units))) @@ -121,7 +122,8 @@ class OpenAIActorCriticConvNet(nn.Module, ActorCriticNet): class OpenAIConvNet(nn.Module, VanillaNet): def __init__(self, in_channels, - n_actions): + n_actions, + gpu=False): super(OpenAIConvNet, self).__init__() self.conv1 = nn.Conv2d(in_channels, 32, 3, stride=2, padding=1) self.conv2 = nn.Conv2d(32, 32, 3, stride=2, padding=1) @@ -132,7 +134,7 @@ class OpenAIConvNet(nn.Module, VanillaNet): self.layer5 = nn.Linear(32 * 3 * 3, hidden_units) self.fc6 = nn.Linear(hidden_units, n_actions) - BasicNet.__init__(self, gpu=False, LSTM=False) + BasicNet.__init__(self, gpu=gpu, LSTM=False) def forward(self, x, update_LSTM=True): x = self.to_torch_variable(x)