[RLlib] ES env_config is not a EnvContext object (e.g. does not contain worker_index). (#7560)

This commit is contained in:
Sven Mika
2020-03-11 20:33:20 -07:00
committed by GitHub
parent bc120730e5
commit dded5b6d22
2 changed files with 15 additions and 11 deletions
+10 -9
View File
@@ -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
+5 -2
View File
@@ -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)