Upgrade DDPG

This commit is contained in:
Shangtong Zhang
2018-05-17 17:17:12 -06:00
parent ab1067d8a8
commit 08038f9533
7 changed files with 109 additions and 40 deletions
+20 -26
View File
@@ -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)
+16
View File
@@ -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
+23 -3
View File
@@ -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()
+8
View File
@@ -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
+33
View File
@@ -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))
-3
View File
@@ -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]
+9 -8
View File
@@ -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()