mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-10 11:40:58 +08:00
Support DM Control
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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')
|
||||
|
||||
Reference in New Issue
Block a user