diff --git a/deep_rl/agent/DDPG_agent.py b/deep_rl/agent/DDPG_agent.py index 6b48e08..bf8f143 100644 --- a/deep_rl/agent/DDPG_agent.py +++ b/deep_rl/agent/DDPG_agent.py @@ -13,15 +13,9 @@ class DDPGAgent(BaseAgent): BaseAgent.__init__(self, config) self.config = config self.task = config.task_fn() - self.network = DisjointActorCriticWrapper(self.task.state_dim, self.task.action_dim, - config.actor_network_fn, config.critic_network_fn) - self.actor = self.network.actor - self.critic = self.network.critic - self.target_network = DisjointActorCriticWrapper(self.task.state_dim, self.task.action_dim, - config.actor_network_fn, config.critic_network_fn) + self.network = config.network_fn(self.task.state_dim, self.task.action_dim) + self.target_network = config.network_fn(self.task.state_dim, self.task.action_dim) self.target_network.load_state_dict(self.network.state_dict()) - self.actor_opt = config.actor_optimizer_fn(self.actor.parameters()) - self.critic_opt = config.critic_optimizer_fn(self.critic.parameters()) self.replay = config.replay_fn() self.random_process = config.random_process_fn(self.task.action_dim) self.total_steps = 0 @@ -35,30 +29,24 @@ class DDPGAgent(BaseAgent): 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() + action = self.network.predict(state, to_numpy=True).flatten() self.config.state_normalizer.unset_read_only() return action - def episode(self, deterministic=False, video_recorder=None): + def episode(self, deterministic=False): self.random_process.reset_states() state = self.task.reset() state = self.config.state_normalizer(state) config = self.config - actor = self.network.actor - critic = self.network.critic - target_actor = self.target_network.actor - target_critic = self.target_network.critic steps = 0 total_reward = 0.0 while True: - action = actor.predict(np.stack([state]), True).flatten() + action = self.network.predict(np.stack([state]), True).flatten() if not deterministic: action += self.random_process.sample() next_state, reward, done, info = self.task.step(action) - if video_recorder is not None: - video_recorder.capture_frame() next_state = self.config.state_normalizer(next_state) total_reward += reward reward = self.config.reward_normalizer(reward) @@ -75,24 +63,30 @@ class DDPGAgent(BaseAgent): if not deterministic and self.replay.size() >= config.min_memory_size: experiences = self.replay.sample() states, actions, rewards, next_states, terminals = experiences - q_next = target_critic.predict(next_states, target_actor.predict(next_states)) - terminals = critic.tensor(terminals).unsqueeze(1) - rewards = critic.tensor(rewards).unsqueeze(1) + + phi_next = self.target_network.feature(next_states) + a_next = self.target_network.actor(phi_next) + q_next = self.target_network.critic(phi_next, a_next) + terminals = self.network.tensor(terminals).unsqueeze(1) + rewards = self.network.tensor(rewards).unsqueeze(1) q_next = config.discount * q_next * (1 - terminals) q_next.add_(rewards) q_next = q_next.detach() - q = critic.predict(states, actions) + phi = self.network.feature(states) + q = self.network.critic(phi, self.network.tensor(actions)) critic_loss = (q - q_next).pow(2).mul(0.5).sum(-1).mean() - self.critic_opt.zero_grad() + self.network.zero_grad() critic_loss.backward() - self.critic_opt.step() + self.network.critic_opt.step() - policy_loss = -critic.predict(states, actor.predict(states)).mean() + phi = self.network.feature(states) + action = self.network.actor(phi) + policy_loss = -self.network.critic(phi.detach(), action).mean() - self.actor_opt.zero_grad() + self.network.zero_grad() policy_loss.backward() - self.actor_opt.step() + self.network.actor_opt.step() self.soft_update(self.target_network, self.network) diff --git a/deep_rl/component/atari_wrapper.py b/deep_rl/component/atari_wrapper.py index 467ccb9..87cda0c 100644 --- a/deep_rl/component/atari_wrapper.py +++ b/deep_rl/component/atari_wrapper.py @@ -254,6 +254,22 @@ class DatasetEnv(gym.Wrapper): self.saved_obs.append(obs) return obs +class RenderEnv(gym.Wrapper): + def __init__(self, env): + gym.Wrapper.__init__(self, env) + self.observation_space = spaces.Box(low=0, high=255, + shape=(self.env.unwrapped._render_height, self.env.unwrapped._render_width, 3), dtype=np.uint8) + + def step(self, action): + _, reward, done, info = self.env.step(action) + obs = self.env.render('rgb_array') + return obs, reward, done, info + + def reset(self): + self.env.reset() + obs = self.env.render('rgb_array') + return obs + def make_atari(env_id, frame_skip=4): env = gym.make(env_id) assert 'NoFrameskip' in env.spec.id diff --git a/deep_rl/component/task.py b/deep_rl/component/task.py index 7acdbc7..20183eb 100644 --- a/deep_rl/component/task.py +++ b/deep_rl/component/task.py @@ -118,6 +118,23 @@ class Bullet(BaseTask): def step(self, action): return BaseTask.step(self, np.clip(action, -1, 1)) +class PixelBullet(BaseTask): + def __init__(self, name, seed=0, log_dir=None, frame_skip=4, history_length=4): + import pybullet_envs + self.name = name + env = gym.make(name) + env.seed(seed) + env = RenderEnv(env) + env = self.set_monitor(env, log_dir) + env = SkipEnv(env, skip=frame_skip) + env = WarpFrame(env) + env = WrapPyTorch(env) + if history_length: + env = StackFrame(env, history_length) + self.action_dim = env.action_space.shape[0] + self.state_dim = env.observation_space.shape + self.env = env + class ProcessTask: def __init__(self, task_fn, log_dir=None): self.pipe, worker_pipe = mp.Pipe() @@ -171,8 +188,11 @@ class ProcessWrapper(mp.Process): raise Exception('Unknown command') class ParallelizedTask: - def __init__(self, task_fn, num_workers, log_dir=None): - self.tasks = [ProcessTask(task_fn, log_dir) for _ in range(num_workers)] + def __init__(self, task_fn, num_workers, log_dir=None, single_process=False): + if single_process: + self.tasks = [task_fn(log_dir=log_dir) for _ in range(num_workers)] + else: + self.tasks = [ProcessTask(task_fn, log_dir) for _ in range(num_workers)] self.state_dim = self.tasks[0].state_dim self.action_dim = self.tasks[0].action_dim self.name = self.tasks[0].name @@ -187,4 +207,4 @@ class ParallelizedTask: return np.stack(results) def close(self): - for task in self.tasks: task.close() + for task in self.tasks: task.close() \ No newline at end of file diff --git a/deep_rl/network/network_bodies.py b/deep_rl/network/network_bodies.py index ad73360..7b706ad 100644 --- a/deep_rl/network/network_bodies.py +++ b/deep_rl/network/network_bodies.py @@ -50,6 +50,14 @@ class TwoLayerFCBodyWithAction(nn.Module): phi = self.gate(self.fc2(torch.cat([x, action], dim=1))) return phi +class DummyBody(nn.Module): + def __init__(self, state_dim): + super(DummyBody, self).__init__() + self.feature_dim = state_dim + + def forward(self, x): + return x + diff --git a/deep_rl/network/network_heads.py b/deep_rl/network/network_heads.py index 33e9f42..3377d4e 100644 --- a/deep_rl/network/network_heads.py +++ b/deep_rl/network/network_heads.py @@ -166,3 +166,36 @@ class DeterministicCriticNet(nn.Module, BaseNet): phi = self.body(x, action) value = self.fc_value(phi) return value + +class DeterministicActorCriticNet(nn.Module, BaseNet): + def __init__(self, action_dim, phi_body, actor_body, critic_body, actor_opt_fn, critic_opt_fn, gpu=-1): + super(DeterministicActorCriticNet, self).__init__() + self.phi_body = phi_body + self.actor_body = actor_body + self.critic_body = critic_body + self.fc_action = layer_init(nn.Linear(actor_body.feature_dim, action_dim), 1e-3) + self.fc_critic = layer_init(nn.Linear(critic_body.feature_dim, 1), 1e-3) + + self.actor_params = list(self.actor_body.parameters()) + list(self.fc_action.parameters()) + self.critic_params = list(self.critic_body.parameters()) + list(self.fc_critic.parameters()) + self.phi_params = list(self.phi_body.parameters()) + self.actor_opt = actor_opt_fn(self.actor_params + self.phi_params) + self.critic_opt = critic_opt_fn(self.critic_params + self.phi_params) + self.set_gpu(gpu) + + def predict(self, obs, to_numpy=False): + phi = self.feature(obs) + action = self.actor(phi) + if to_numpy: + return action.cpu().detach().numpy() + return action + + def feature(self, obs): + obs = self.tensor(obs) + return self.phi_body(obs) + + def actor(self, phi): + return F.tanh(self.fc_action(self.actor_body(phi))) + + def critic(self, phi, a): + return self.fc_critic(self.critic_body(phi, a)) diff --git a/deep_rl/utils/misc.py b/deep_rl/utils/misc.py index 8d21b5f..a4aa07c 100644 --- a/deep_rl/utils/misc.py +++ b/deep_rl/utils/misc.py @@ -120,6 +120,3 @@ class Batcher: indices = np.arange(self.num_entries) np.random.shuffle(indices) self.data = [d[indices] for d in self.data] - -# def torch_max(tensor, dim): -# return torch.max(tensor, dim=dim, keepdim=True)[0] diff --git a/examples.py b/examples.py index b4fb7be..523eedf 100644 --- a/examples.py +++ b/examples.py @@ -339,7 +339,7 @@ def ppo_continuous(): def ddpg_continuous(): config = Config() log_dir = get_default_log_dir(ddpg_continuous.__name__) - # config.task_fn = lambda: Pendulum(log_dir=log_dir) + # task_fn = lambda **kwargs: Pendulum(log_dir=log_dir) task_fn = lambda **kwargs: Bullet('AntBulletEnv-v0', **kwargs) # each bullet environment should be started in a new process, it is a workaround @@ -348,12 +348,13 @@ def ddpg_continuous(): config.task_fn = lambda: ProcessTask(task_fn) config.evaluation_env = ProcessTask(task_fn, log_dir=log_dir) - config.actor_network_fn = lambda state_dim, action_dim: DeterministicActorNet( - action_dim, FCBody(state_dim, (300, 200), gate=F.tanh)) - config.critic_network_fn = lambda state_dim, action_dim: DeterministicCriticNet( - TwoLayerFCBodyWithAction(state_dim, action_dim, (400, 300), gate=F.tanh)) - config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4) - config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-3) + config.network_fn = lambda state_dim, action_dim: DeterministicActorCriticNet( + action_dim=action_dim, phi_body=DummyBody(state_dim), + actor_body=FCBody(state_dim, (300, 200), gate=F.tanh), + critic_body=TwoLayerFCBodyWithAction(state_dim, action_dim, (400, 300), gate=F.tanh), + actor_opt_fn=lambda params: torch.optim.Adam(params, lr=1e-4), + critic_opt_fn=lambda params: torch.optim.Adam(params, lr=1e-3)) + config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=64) config.discount = 0.99 config.state_normalizer = RunningStatsNormalizer() @@ -415,7 +416,7 @@ if __name__ == '__main__': # option_ciritc_pixel_atari('BreakoutNoFrameskip-v4') # dqn_ram_atari('Breakout-ramNoFrameskip-v4') - # ddpg_continuous() + ddpg_continuous() # ppo_continuous() # action_conditional_video_prediction()