mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Upgrade DDPG
This commit is contained in:
+20
-26
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user