Upgrade all actor critic methods

This commit is contained in:
Shangtong Zhang
2018-05-17 23:37:17 -06:00
parent aa07d467bd
commit 84f7911bb4
5 changed files with 207 additions and 220 deletions
+11 -16
View File
@@ -15,7 +15,6 @@ class A2CAgent(BaseAgent):
self.task = config.task_fn()
self.network = config.network_fn(self.task.state_dim, self.task.action_dim)
self.optimizer = config.optimizer_fn(self.network.parameters())
self.policy = config.policy_fn()
self.total_steps = 0
self.states = self.task.reset()
self.episode_rewards = np.zeros(config.num_workers)
@@ -26,9 +25,8 @@ class A2CAgent(BaseAgent):
rollout = []
states = self.states
for _ in range(config.rollout_length):
prob, log_prob, value = self.network.predict(config.state_normalizer(states))
actions = [self.policy.sample(p) for p in prob.cpu().detach().numpy()]
next_states, rewards, terminals, _ = self.task.step(actions)
actions, log_probs, entropy, values = self.network.predict(config.state_normalizer(states))
next_states, rewards, terminals, _ = self.task.step(actions.detach().cpu().numpy())
self.episode_rewards += rewards
rewards = config.reward_normalizer(rewards)
for i, terminal in enumerate(terminals):
@@ -36,21 +34,20 @@ class A2CAgent(BaseAgent):
self.last_episode_rewards[i] = self.episode_rewards[i]
self.episode_rewards[i] = 0
rollout.append([prob, log_prob, value, actions, rewards, 1 - terminals])
rollout.append([log_probs, values, actions, rewards, 1 - terminals, entropy])
states = next_states
self.states = states
_, _, pending_value = self.network.predict(config.state_normalizer(states))
rollout.append([None, None, pending_value, None, None, None])
pending_value = self.network.predict(config.state_normalizer(states))[-1]
rollout.append([None, pending_value, None, None, None, None])
processed_rollout = [None] * (len(rollout) - 1)
advantages = self.network.tensor(np.zeros((config.num_workers, 1)))
returns = pending_value.detach()
for i in reversed(range(len(rollout) - 1)):
prob, log_prob, value, actions, rewards, terminals = rollout[i]
log_prob, value, actions, rewards, terminals, entropy = rollout[i]
terminals = self.network.tensor(terminals).unsqueeze(1)
rewards = self.network.tensor(rewards).unsqueeze(1)
actions = self.network.tensor(actions).unsqueeze(1).long()
next_value = rollout[i + 1][2]
returns = rewards + config.discount * terminals * returns
if not config.use_gae:
@@ -58,24 +55,22 @@ class A2CAgent(BaseAgent):
else:
td_error = rewards + config.discount * terminals * next_value.detach() - value.detach()
advantages = advantages * config.gae_tau * config.discount * terminals + td_error
processed_rollout[i] = [prob, log_prob, value, actions, returns, advantages]
processed_rollout[i] = [log_prob, value, returns, advantages, entropy]
prob, log_prob, value, actions, returns, advantages = map(lambda x: torch.cat(x, dim=0), zip(*processed_rollout))
policy_loss = -log_prob.gather(1, actions) * advantages
entropy_loss = torch.sum(prob * log_prob, dim=1, keepdim=True)
log_prob, value, returns, advantages, entropy = map(lambda x: torch.cat(x, dim=0), zip(*processed_rollout))
policy_loss = -log_prob * advantages
value_loss = 0.5 * (returns - value).pow(2)
entropy_loss = entropy.mean()
self.policy_loss = np.mean(policy_loss.cpu().detach().numpy())
self.entropy_loss = np.mean(entropy_loss.cpu().detach().numpy())
self.value_loss = np.mean(value_loss.cpu().detach().numpy())
self.optimizer.zero_grad()
(policy_loss + config.entropy_weight * entropy_loss +
(policy_loss - config.entropy_weight * entropy_loss +
config.value_loss_weight * value_loss).mean().backward()
nn.utils.clip_grad_norm_(self.network.parameters(), config.gradient_clip)
self.optimizer.step()
self.evaluate(config.rollout_length)
steps = config.rollout_length * config.num_workers
self.total_steps += steps
+4 -3
View File
@@ -14,6 +14,7 @@ class PPOAgent(BaseAgent):
self.config = config
self.task = config.task_fn()
self.network = config.network_fn(self.task.state_dim, self.task.action_dim)
self.opt = config.optimizer_fn(self.network.parameters())
self.total_steps = 0
self.episode_rewards = np.zeros(config.num_workers)
self.last_episode_rewards = np.zeros(config.num_workers)
@@ -79,14 +80,14 @@ class PPOAgent(BaseAgent):
obj = ratio * sampled_advantages
obj_clipped = ratio.clamp(1.0 - self.config.ppo_ratio_clip,
1.0 + self.config.ppo_ratio_clip) * sampled_advantages
policy_loss = -torch.min(obj, obj_clipped).mean(0) + config.entropy_weight * entropy_loss
policy_loss = -torch.min(obj, obj_clipped).mean(0) - config.entropy_weight * entropy_loss.mean()
value_loss = 0.5 * (sampled_returns - values).pow(2).mean()
self.network.zero_grad()
self.opt.zero_grad()
(policy_loss + value_loss).backward()
nn.utils.clip_grad_norm_(self.network.parameters(), config.gradient_clip)
self.network.step()
self.opt.step()
steps = config.rollout_length * config.num_workers
self.total_steps += steps
+79 -84
View File
@@ -5,6 +5,7 @@
#######################################################################
from .network_utils import *
from .network_bodies import *
class VanillaNet(nn.Module, BaseNet):
def __init__(self, output_dim, body, gpu=-1):
@@ -37,24 +38,6 @@ class DuelingNet(nn.Module, BaseNet):
return q.cpu().detach().numpy()
return q
class ActorCriticNet(nn.Module, BaseNet):
def __init__(self, action_dim, body, gpu=-1):
super(ActorCriticNet, self).__init__()
self.fc_actor = layer_init(nn.Linear(body.feature_dim, action_dim))
self.fc_critic = layer_init(nn.Linear(body.feature_dim, 1))
self.body = body
self.set_gpu(gpu)
def predict(self, x, to_numpy=False):
phi = self.body(self.tensor(x))
pre_prob = self.fc_actor(phi)
prob = F.softmax(pre_prob, dim=1)
log_prob = F.log_softmax(pre_prob, dim=1)
value = self.fc_critic(phi)
if to_numpy:
return prob.cpu().detach().numpy()
return prob, log_prob, value
class CategoricalNet(nn.Module, BaseNet):
def __init__(self, action_dim, num_atoms, body, gpu=-1):
super(CategoricalNet, self).__init__()
@@ -109,67 +92,12 @@ class OptionCriticNet(nn.Module, BaseNet):
log_pi = F.log_softmax(pi, dim=-1)
return q, beta, log_pi
class GaussianActorNet(nn.Module, BaseNet):
def __init__(self, action_dim, body, gpu=-1):
super(GaussianActorNet, self).__init__()
self.fc_action = layer_init(nn.Linear(body.feature_dim, action_dim), 3e-3)
self.action_log_std = nn.Parameter(torch.zeros(1, action_dim))
self.body = body
self.set_gpu(gpu)
def predict(self, x):
x = self.tensor(x)
phi = self.body(x)
mean = F.tanh(self.fc_action(phi))
log_std = self.action_log_std.expand_as(mean)
std = log_std.exp()
return mean, std, log_std
class GaussianCriticNet(nn.Module, BaseNet):
def __init__(self, body, gpu=-1):
super(GaussianCriticNet, self).__init__()
self.fc_value = layer_init(nn.Linear(body.feature_dim, 1), 3e-3)
self.body = body
self.set_gpu(gpu)
def predict(self, x):
x = self.tensor(x)
phi = self.body(x)
value = self.fc_value(phi)
return value
class DeterministicActorNet(nn.Module, BaseNet):
def __init__(self, action_dim, body, gpu=-1):
super(DeterministicActorNet, self).__init__()
self.fc_action = layer_init(nn.Linear(body.feature_dim, action_dim), 3e-3)
self.body = body
self.set_gpu(gpu)
def predict(self, x, to_numpy=False):
x = self.tensor(x)
phi = self.body(x)
a = F.tanh(self.fc_action(phi))
if to_numpy:
a = a.cpu().detach().numpy()
return a
class DeterministicCriticNet(nn.Module, BaseNet):
def __init__(self, body, gpu=-1):
super(DeterministicCriticNet, self).__init__()
self.fc_value = layer_init(nn.Linear(body.feature_dim, 1), 3e-3)
self.body = body
self.set_gpu(gpu)
def predict(self, x, action):
x = self.tensor(x)
action = self.tensor(action)
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__()
class ActorCriticNet(nn.Module):
def __init__(self, state_dim, action_dim, phi_body, actor_body, critic_body):
super(ActorCriticNet, self).__init__()
if phi_body is None: phi_body = DummyBody(state_dim)
if actor_body is None: actor_body = DummyBody(phi_body.feature_dim)
if critic_body is None: critic_body = DummyBody(phi_body.feature_dim)
self.phi_body = phi_body
self.actor_body = actor_body
self.critic_body = critic_body
@@ -179,8 +107,21 @@ class DeterministicActorCriticNet(nn.Module, BaseNet):
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)
class DeterministicActorCriticNet(nn.Module, BaseNet):
def __init__(self,
state_dim,
action_dim,
actor_opt_fn,
critic_opt_fn,
phi_body=None,
actor_body=None,
critic_body=None,
gpu=-1):
super(DeterministicActorCriticNet, self).__init__()
self.network = ActorCriticNet(state_dim, action_dim, phi_body, actor_body, critic_body)
self.actor_opt = actor_opt_fn(self.network.actor_params + self.network.phi_params)
self.critic_opt = critic_opt_fn(self.network.critic_params + self.network.phi_params)
self.set_gpu(gpu)
def predict(self, obs, to_numpy=False):
@@ -192,10 +133,64 @@ class DeterministicActorCriticNet(nn.Module, BaseNet):
def feature(self, obs):
obs = self.tensor(obs)
return self.phi_body(obs)
return self.network.phi_body(obs)
def actor(self, phi):
return F.tanh(self.fc_action(self.actor_body(phi)))
return F.tanh(self.network.fc_action(self.network.actor_body(phi)))
def critic(self, phi, a):
return self.fc_critic(self.critic_body(phi, a))
return self.network.fc_critic(self.network.critic_body(phi, a))
class GaussianActorCriticNet(nn.Module, BaseNet):
def __init__(self,
state_dim,
action_dim,
phi_body=None,
actor_body=None,
critic_body=None,
gpu=-1):
super(GaussianActorCriticNet, self).__init__()
self.network = ActorCriticNet(state_dim, action_dim, phi_body, actor_body, critic_body)
self.std = nn.Parameter(torch.ones(1, action_dim))
self.set_gpu(gpu)
def predict(self, obs, action=None, to_numpy=False):
obs = self.tensor(obs)
phi = self.network.phi_body(obs)
phi_a = self.network.actor_body(phi)
phi_v = self.network.critic_body(phi)
mean = F.tanh(self.network.fc_action(phi_a))
if to_numpy:
return mean.cpu().detach().numpy()
v = self.network.fc_critic(phi_v)
dist = torch.distributions.Normal(mean, self.std)
if action is None:
action = dist.sample()
log_prob = dist.log_prob(action)
log_prob = torch.sum(log_prob, dim=1, keepdim=True)
return action, log_prob, self.tensor(np.zeros((log_prob.size(0), 1))), v
class CategoricalActorCriticNet(nn.Module, BaseNet):
def __init__(self,
state_dim,
action_dim,
phi_body=None,
actor_body=None,
critic_body=None,
gpu=-1):
super(CategoricalActorCriticNet, self).__init__()
self.network = ActorCriticNet(state_dim, action_dim, phi_body, actor_body, critic_body)
self.set_gpu(gpu)
def predict(self, obs, action=None):
obs = self.tensor(obs)
phi = self.network.phi_body(obs)
phi_a = self.network.actor_body(phi)
phi_v = self.network.critic_body(phi)
prob = F.softmax(self.network.fc_action(phi_a), dim=-1)
v = self.network.fc_critic(phi_v)
dist = torch.distributions.Categorical(probs=prob)
if action is None:
action = dist.sample()
log_prob = dist.log_prob(action).unsqueeze(-1)
return action, log_prob, dist.entropy().unsqueeze(-1), v
+88 -88
View File
@@ -23,94 +23,94 @@ class BaseNet:
x = torch.tensor(x, device=self.device, dtype=torch.float32)
return x
class DisjointActorCriticWrapper:
def __init__(self, state_dim, action_dim, actor_network_fn, critic_network_fn):
self.actor = actor_network_fn(state_dim, action_dim)
self.critic = critic_network_fn(state_dim, action_dim)
def state_dict(self):
return [self.actor.state_dict(), self.critic.state_dict()]
def load_state_dict(self, state_dicts):
self.actor.load_state_dict(state_dicts[0])
self.critic.load_state_dict(state_dicts[1])
def parameters(self):
return list(self.actor.parameters()) + list(self.critic.parameters())
def zero_grad(self):
self.actor.zero_grad()
self.critic.zero_grad()
class GaussianActorCriticWrapper:
def __init__(self, state_dim, action_dim, actor_fn, critic_fn, actor_opt_fn, critic_opt_fn):
self.actor = actor_fn(state_dim, action_dim)
self.critic = critic_fn(state_dim)
self.actor_opt = actor_opt_fn(self.actor.parameters())
self.critic_opt = critic_opt_fn(self.critic.parameters())
def predict(self, state, actions=None):
mean, std, log_std = self.actor.predict(state)
values = self.critic.predict(state)
dist = torch.distributions.Normal(mean, std)
if actions is None:
actions = dist.sample()
log_probs = dist.log_prob(actions)
log_probs = torch.sum(log_probs, dim=1, keepdim=True)
return actions, log_probs, 0, values
def tensor(self, x):
return self.actor.tensor(x)
def zero_grad(self):
self.actor_opt.zero_grad()
self.critic_opt.zero_grad()
def parameters(self):
return list(self.actor.parameters()) + list(self.critic.parameters())
def step(self):
self.actor_opt.step()
self.critic_opt.step()
def state_dict(self):
return [self.actor.state_dict(), self.critic.state_dict()]
def load_state_dict(self, state_dicts):
self.actor.load_state_dict(state_dicts[0])
self.critic.load_state_dict(state_dicts[1])
class CategoricalActorCriticWrapper:
def __init__(self, state_dim, action_dim, network_fn, opt_fn):
self.network = network_fn(state_dim, action_dim)
self.opt = opt_fn(self.network.parameters())
def predict(self, state, action=None):
prob, log_prob, value = self.network.predict(state)
entropy_loss = torch.sum(prob * log_prob, dim=1, keepdim=True)
dist = torch.distributions.Categorical(prob)
if action is None:
action = dist.sample()
log_prob = dist.log_prob(action).unsqueeze(1)
return action, log_prob, entropy_loss.mean(0), value
def tensor(self, x):
return self.network.tensor(x)
def zero_grad(self):
self.opt.zero_grad()
def parameters(self):
return self.network.parameters()
def step(self):
self.opt.step()
def state_dict(self):
return self.network.state_dict()
def load_state_dict(self, state_dicts):
self.network.load_state_dict(state_dicts)
# class DisjointActorCriticWrapper:
# def __init__(self, state_dim, action_dim, actor_network_fn, critic_network_fn):
# self.actor = actor_network_fn(state_dim, action_dim)
# self.critic = critic_network_fn(state_dim, action_dim)
#
# def state_dict(self):
# return [self.actor.state_dict(), self.critic.state_dict()]
#
# def load_state_dict(self, state_dicts):
# self.actor.load_state_dict(state_dicts[0])
# self.critic.load_state_dict(state_dicts[1])
#
# def parameters(self):
# return list(self.actor.parameters()) + list(self.critic.parameters())
#
# def zero_grad(self):
# self.actor.zero_grad()
# self.critic.zero_grad()
#
# class GaussianActorCriticWrapper:
# def __init__(self, state_dim, action_dim, actor_fn, critic_fn, actor_opt_fn, critic_opt_fn):
# self.actor = actor_fn(state_dim, action_dim)
# self.critic = critic_fn(state_dim)
# self.actor_opt = actor_opt_fn(self.actor.parameters())
# self.critic_opt = critic_opt_fn(self.critic.parameters())
#
# def predict(self, state, actions=None):
# mean, std, log_std = self.actor.predict(state)
# values = self.critic.predict(state)
# dist = torch.distributions.Normal(mean, std)
# if actions is None:
# actions = dist.sample()
# log_probs = dist.log_prob(actions)
# log_probs = torch.sum(log_probs, dim=1, keepdim=True)
# return actions, log_probs, 0, values
#
# def tensor(self, x):
# return self.actor.tensor(x)
#
# def zero_grad(self):
# self.actor_opt.zero_grad()
# self.critic_opt.zero_grad()
#
# def parameters(self):
# return list(self.actor.parameters()) + list(self.critic.parameters())
#
# def step(self):
# self.actor_opt.step()
# self.critic_opt.step()
#
# def state_dict(self):
# return [self.actor.state_dict(), self.critic.state_dict()]
#
# def load_state_dict(self, state_dicts):
# self.actor.load_state_dict(state_dicts[0])
# self.critic.load_state_dict(state_dicts[1])
#
# class CategoricalActorCriticWrapper:
# def __init__(self, state_dim, action_dim, network_fn, opt_fn):
# self.network = network_fn(state_dim, action_dim)
# self.opt = opt_fn(self.network.parameters())
#
# def predict(self, state, action=None):
# prob, log_prob, value = self.network.predict(state)
# entropy_loss = torch.sum(prob * log_prob, dim=1, keepdim=True)
# dist = torch.distributions.Categorical(prob)
# if action is None:
# action = dist.sample()
# log_prob = dist.log_prob(action).unsqueeze(1)
# return action, log_prob, entropy_loss.mean(0), value
#
# def tensor(self, x):
# return self.network.tensor(x)
#
# def zero_grad(self):
# self.opt.zero_grad()
#
# def parameters(self):
# return self.network.parameters()
#
# def step(self):
# self.opt.step()
#
# def state_dict(self):
# return self.network.state_dict()
#
# def load_state_dict(self, state_dicts):
# self.network.load_state_dict(state_dicts)
def layer_init(layer, w_scale=1.0):
nn.init.orthogonal_(layer.weight.data)
+25 -29
View File
@@ -31,12 +31,12 @@ def a2c_cart_pole():
name = 'CartPole-v0'
# name = 'MountainCar-v0'
task_fn = lambda log_dir: ClassicalControl(name, max_steps=200, log_dir=log_dir)
config.evaluation_env = task_fn(None)
config.num_workers = 5
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers,
log_dir=get_default_log_dir(a2c_cart_pole.__name__))
config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
config.network_fn = lambda state_dim, action_dim: ActorCriticNet(action_dim, FCBody(state_dim))
config.network_fn = lambda state_dim, action_dim: CategoricalActorCriticNet(
state_dim, action_dim, FCBody(state_dim), gpu=-1)
config.policy_fn = SamplePolicy
config.discount = 0.99
config.logger = get_logger()
@@ -100,10 +100,9 @@ def ppo_cart_pole():
task_fn = lambda log_dir: ClassicalControl('CartPole-v0', max_steps=200, log_dir=log_dir)
config.num_workers = 5
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001)
network_fn = lambda state_dim, action_dim: ActorCriticNet(action_dim, FCBody(state_dim))
config.network_fn = lambda state_dim, action_dim: \
CategoricalActorCriticWrapper(state_dim, action_dim, network_fn, optimizer_fn)
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001)
config.network_fn = lambda state_dim, action_dim: CategoricalActorCriticNet(
state_dim, action_dim, FCBody(state_dim), gpu=-1)
config.discount = 0.99
config.logger = get_logger()
config.use_gae = True
@@ -164,8 +163,8 @@ def a2c_pixel_atari(name):
task_fn = lambda log_dir: PixelAtari(name, frame_skip=4, history_length=config.history_length, log_dir=log_dir)
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, log_dir=get_default_log_dir(a2c_pixel_atari.__name__))
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.0007)
config.network_fn = lambda state_dim, action_dim: \
ActorCriticNet(action_dim, NatureConvBody(), gpu=1)
config.network_fn = lambda state_dim, action_dim: CategoricalActorCriticNet(
state_dim, action_dim, NatureConvBody(), gpu=0)
config.policy_fn = SamplePolicy
config.state_normalizer = ImageNormalizer()
config.reward_normalizer = SignNormalizer()
@@ -246,10 +245,9 @@ def ppo_pixel_atari(name):
config.num_workers = 16
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers,
log_dir=get_default_log_dir(ppo_pixel_atari.__name__))
optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025)
network_fn = lambda state_dim, action_dim: ActorCriticNet(action_dim, NatureConvBody(), gpu=2)
config.network_fn = lambda state_dim, action_dim: \
CategoricalActorCriticWrapper(state_dim, action_dim, network_fn, optimizer_fn)
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025)
config.network_fn = lambda state_dim, action_dim: CategoricalActorCriticNet(
state_dim, action_dim, NatureConvBody(), gpu=0)
config.state_normalizer = ImageNormalizer()
config.reward_normalizer = SignNormalizer()
config.discount = 0.99
@@ -312,17 +310,14 @@ def ppo_continuous():
config = Config()
config.num_workers = 1
# task_fn = lambda log_dir: Pendulum(log_dir=log_dir)
task_fn = lambda log_dir: Bullet('AntBulletEnv-v0', log_dir=log_dir)
# task_fn = lambda log_dir: Bullet('AntBulletEnv-v0', log_dir=log_dir)
task_fn = lambda log_dir: Roboschool('RoboschoolAnt-v1', log_dir=log_dir)
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, log_dir=get_default_log_dir(ppo_continuous.__name__))
actor_network_fn = lambda state_dim, action_dim: GaussianActorNet(
action_dim, FCBody(state_dim))
critic_network_fn = lambda state_dim: GaussianCriticNet(FCBody(state_dim))
actor_optimizer_fn = lambda params: torch.optim.Adam(params, 3e-4, eps=1e-5)
critic_optimizer_fn = lambda params: torch.optim.Adam(params, 3e-4, eps=1e-5)
config.network_fn = lambda state_dim, action_dim: \
GaussianActorCriticWrapper(state_dim, action_dim, actor_network_fn,
critic_network_fn, actor_optimizer_fn,
critic_optimizer_fn)
config.network_fn = lambda state_dim, action_dim: GaussianActorCriticNet(
state_dim, action_dim, actor_body=FCBody(state_dim),
critic_body=FCBody(state_dim), gpu=-1)
config.optimizer_fn = lambda params: torch.optim.Adam(params, 3e-4, eps=1e-5)
# config.state_normalizer = RunningStatsNormalizer()
config.discount = 0.99
config.use_gae = True
@@ -336,11 +331,12 @@ def ppo_continuous():
config.logger = get_logger()
run_iterations(PPOAgent(config))
def ddpg_internal_state():
def ddpg_low_dim_state():
config = Config()
log_dir = get_default_log_dir(ddpg_internal_state.__name__)
log_dir = get_default_log_dir(ddpg_low_dim_state.__name__)
# task_fn = lambda **kwargs: Pendulum(log_dir=log_dir)
task_fn = lambda **kwargs: Bullet('AntBulletEnv-v0', **kwargs)
# task_fn = lambda **kwargs: Bullet('AntBulletEnv-v0', **kwargs)
task_fn = lambda **kwargs: Roboschool('RoboschoolAnt-v1', **kwargs)
# each bullet environment should be started in a new process, it is a workaround
# to the issue of self-collision
@@ -349,7 +345,7 @@ def ddpg_internal_state():
config.evaluation_env = ProcessTask(task_fn, log_dir=log_dir)
config.network_fn = lambda state_dim, action_dim: DeterministicActorCriticNet(
action_dim=action_dim, phi_body=DummyBody(state_dim),
state_dim, action_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),
@@ -377,7 +373,7 @@ def ddpg_pixel():
phi_body=NatureConvBody()
config.network_fn = lambda state_dim, action_dim: DeterministicActorCriticNet(
action_dim=action_dim, phi_body=NatureConvBody(),
state_dim, action_dim, phi_body=NatureConvBody(),
actor_body=FCBody(phi_body.feature_dim, (200, 200), gate=F.relu),
critic_body=TwoLayerFCBodyWithAction(phi_body.feature_dim, action_dim, (200, 200), gate=F.relu),
actor_opt_fn=lambda params: torch.optim.Adam(params, lr=1e-4),
@@ -446,8 +442,8 @@ if __name__ == '__main__':
# option_ciritc_pixel_atari('BreakoutNoFrameskip-v4')
# dqn_ram_atari('Breakout-ramNoFrameskip-v4')
# ddpg_internal_state()
ddpg_pixel()
# ddpg_low_dim_state()
# ddpg_pixel()
# ppo_continuous()
# action_conditional_video_prediction()