From dded5b6d22e701a56461796f70e569802545b323 Mon Sep 17 00:00:00 2001 From: Sven Mika Date: Thu, 12 Mar 2020 04:33:20 +0100 Subject: [PATCH] [RLlib] ES `env_config` is not a EnvContext object (e.g. does not contain `worker_index`). (#7560) --- rllib/agents/es/es.py | 19 ++++++++++--------- rllib/agents/es/policies.py | 7 +++++-- 2 files changed, 15 insertions(+), 11 deletions(-) diff --git a/rllib/agents/es/es.py b/rllib/agents/es/es.py index bb152f59a..ff8932f3c 100644 --- a/rllib/agents/es/es.py +++ b/rllib/agents/es/es.py @@ -8,14 +8,12 @@ import time import ray from ray.rllib.agents import Trainer, with_common_config - -from ray.rllib.agents.es import optimizers -from ray.rllib.agents.es import policies -from ray.rllib.agents.es import utils +from ray.rllib.agents.es import optimizers, policies, utils +from ray.rllib.env.env_context import EnvContext from ray.rllib.policy.sample_batch import DEFAULT_POLICY_ID +from ray.rllib.utils import FilterManager from ray.rllib.utils.annotations import override from ray.rllib.utils.memory import ray_get_and_free -from ray.rllib.utils import FilterManager logger = logging.getLogger(__name__) @@ -70,13 +68,15 @@ class Worker: policy_params, env_creator, noise, + worker_index, min_task_runtime=0.2): self.min_task_runtime = min_task_runtime self.config = config self.policy_params = policy_params self.noise = SharedNoiseTable(noise) - self.env = env_creator(config["env_config"]) + env_context = EnvContext(config["env_config"] or {}, worker_index) + self.env = env_creator(env_context) from ray.rllib import models self.preprocessor = models.ModelCatalog.get_preprocessor( self.env, config["model"]) @@ -175,7 +175,8 @@ class ESTrainer(Trainer): policy_params = {"action_noise_std": 0.01} - env = env_creator(config["env_config"]) + env_context = EnvContext(config["env_config"] or {}, worker_index=0) + env = env_creator(env_context) from ray.rllib import models preprocessor = models.ModelCatalog.get_preprocessor(env) @@ -194,8 +195,8 @@ class ESTrainer(Trainer): # Create the actors. logger.info("Creating actors.") self._workers = [ - Worker.remote(config, policy_params, env_creator, noise_id) - for _ in range(config["num_workers"]) + Worker.remote(config, policy_params, env_creator, noise_id, + idx + 1) for idx in range(config["num_workers"]) ] self.episodes_so_far = 0 diff --git a/rllib/agents/es/policies.py b/rllib/agents/es/policies.py index c4a80f9ce..64d0fa67d 100644 --- a/rllib/agents/es/policies.py +++ b/rllib/agents/es/policies.py @@ -20,13 +20,16 @@ def rollout(policy, env, timestep_limit=None, add_noise=False): If add_noise is True, the rollout will take noisy actions with noise drawn from that stream. Otherwise, no action noise will be added. """ - env_timestep_limit = env.spec.max_episode_steps + max_timestep_limit = 999999 + env_timestep_limit = env.spec.max_episode_steps if ( + hasattr(env, "spec") and hasattr(env.spec, "max_episode_steps")) \ + else max_timestep_limit timestep_limit = (env_timestep_limit if timestep_limit is None else min( timestep_limit, env_timestep_limit)) rews = [] t = 0 observation = env.reset() - for _ in range(timestep_limit or 999999): + for _ in range(timestep_limit or max_timestep_limit): ac = policy.compute(observation, add_noise=add_noise)[0] observation, rew, done, _ = env.step(ac) rews.append(rew)