Support DM Control

This commit is contained in:
Shangtong Zhang
2018-04-07 20:53:29 -06:00
parent 8623fe849b
commit 0f3f49252b
3 changed files with 42 additions and 19 deletions
+1
View File
@@ -56,6 +56,7 @@ Prediction is sampled after 110K iterations and I only implemented one-step trai
> Tested in macOS 10.12 and CentO/S 6.8
* OpenAI gym
* [Roboschool](https://github.com/openai/roboschool) (Optional)
* [DeepMind Control Suite](https://github.com/deepmind/dm_control) & [DMControl2Gym](dm_control2gym) (Optional)
* PyTorch v0.3.0
* Python 2.7 / 3.6
* [TensorboardX](https://github.com/lanpa/tensorboard-pytorch)
+15
View File
@@ -139,6 +139,21 @@ class Roboschool(BasicTask):
def step(self, action):
return BasicTask.step(self, np.clip(action, -1, 1))
class DMControl(BasicTask):
def __init__(self, domain_name, task_name, max_steps=sys.maxsize, log_dir=None):
from dm_control import suite
import dm_control2gym
BasicTask.__init__(self, max_steps)
self.name = domain_name + '_' + task_name
self.env = dm_control2gym.make(domain_name, task_name)
self.action_dim = self.env.action_space.shape[0]
self.state_dim = self.env.observation_space.shape[0]
if log_dir is not None:
mkdir(log_dir)
self.env = Monitor(self.env, '%s/%s' % (log_dir, uuid.uuid1()))
def sub_task(parent_pipe, pipe, task_fn, rank, log_dir):
np.random.seed()
seed = np.random.randint(0, sys.maxsize)
+26 -19
View File
@@ -187,7 +187,7 @@ def n_step_dqn_pixel_atari(name):
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers,
log_dir=get_default_log_dir(n_step_dqn_pixel_atari.__name__))
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01)
config.network_fn = lambda state_dim, action_dim: ConvNet(config.history_length, action_dim, gpu=0)
config.network_fn = lambda state_dim, action_dim: ConvNet(config.history_length, action_dim, gpu=3)
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
config.state_normalizer = ImageNormalizer()
config.reward_normalizer = SignNormalizer()
@@ -220,22 +220,24 @@ def dqn_ram_atari(name):
def ppo_continuous():
config = Config()
config.num_workers = 5
task_fn = lambda log_dir: Pendulum(log_dir=log_dir)
config.num_workers = 16
# task_fn = lambda log_dir: Pendulum(log_dir=log_dir)
# task_fn = lambda log_dir: Roboschool('RoboschoolInvertedPendulum-v1', log_dir=log_dir)
# task_fn = lambda log_dir: Roboschool('RoboschoolAnt-v1', log_dir=log_dir)
# task_fn = lambda log_dir: Roboschool('RoboschoolReacher-v1', log_dir=log_dir)
task_fn = lambda log_dir: Roboschool('RoboschoolHopper-v1', log_dir=log_dir)
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, log_dir=get_default_log_dir(ppo_continuous.__name__))
config.actor_network_fn = lambda state_dim, action_dim: GaussianActorNet(state_dim, action_dim)
config.critic_network_fn = lambda state_dim, action_dim: GaussianCriticNet(state_dim)
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 1e-4)
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 1e-4)
config.state_normalizer = RunningStatsNormalizer()
config.discount = 0.99
config.use_gae = True
config.gae_tau = 0.97
config.gradient_clip = 0.5
config.rollout_length = 20
config.optimize_epochs = 4
config.optimize_epochs = 5
config.ppo_ratio_clip = 0.2
config.logger = Logger('./log', logger)
run_iterations(PPOAgent(config))
@@ -243,12 +245,13 @@ 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)
# config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-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: Pendulum(log_dir=log_dir)
# config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1', log_dir=log_dir)
config.task_fn = lambda: Roboschool('RoboschoolReacher-v1', log_dir=log_dir)
# config.task_fn = lambda: Roboschool('RoboschoolHopper-v1', log_dir=log_dir)
# config.task_fn = lambda: Roboschool('RoboschoolAnt-v1', log_dir=log_dir)
# config.task_fn = lambda: Roboschool('RoboschoolWalker2d-v1', log_dir=log_dir)
config.task_fn = lambda: DMControl('cartpole', 'balance', log_dir=log_dir)
config.actor_network_fn = lambda state_dim, action_dim: DeterministicActorNet(state_dim, action_dim)
config.critic_network_fn = lambda state_dim, action_dim: DeterministicCriticNet(state_dim, action_dim)
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4)
@@ -257,7 +260,7 @@ def ddpg_continuous():
config.discount = 0.99
config.random_process_fn = \
lambda action_dim: OrnsteinUhlenbeckProcess(size=action_dim, theta=0.15, sigma=0.3,
n_steps_annealing=100000)
n_steps_annealing=1000000)
config.min_memory_size = 64
config.target_network_mix = 1e-3
config.gradient_clip = 1.0
@@ -267,12 +270,16 @@ def ddpg_continuous():
def plot():
import matplotlib.pyplot as plt
plotter = Plotter()
name = 'to_plot/a2c_pixel_atari-180407-92711'
# name = 'to_plot/dqn_pixel_atari-180407-01414'
# name = 'to_plot/quantile_regression_dqn_pixel_atari-180407-01604'
# name = 'to_plot/categorical_dqn_pixel_atari-180407-01537'
plotter.plot_results([name])
plt.show()
names = ['a2c_pixel_atari-180407-92711',
'categorical_dqn_pixel_atari-180407-094006',
'dqn_pixel_atari-180407-01414',
'quantile_regression_dqn_pixel_atari-180407-01604',
'n_step_dqn_pixel_atari-180407-163421',
'ppo_continuous-180407-111715']
for name in names:
plotter.plot_results(['to_plot/%s' % (name)])
plt.savefig('images/%s.png' % (name))
plt.close()
if __name__ == '__main__':
mkdir('data')