[rllib] Part 1 of multiagent support: make sampler path support multiagent envs (#2268)

This refactors the RLlib sampler to support multi-agent environments. The main changes were:

AsyncVectorEnv now produces dicts of env_id -> agent_id -> value rather than env_id -> value. This lets it model both vectorized and multi-agent envs (or both).
The sampler class operates over the above nested dict structure for all envs. Single agent envs just return a dict with one agent_id=single_agent.
When sample() is called on a policy evaluator, in the single agent case we return a SampleBatch, otherwise we return a MultiAgentBatch (which is a list of sample batches per policy).
Left for another PR:

Exposing multi-agent in the public interfaces.
Optimizations such as evaluating multiple policies in one TF run.
This commit is contained in:
Eric Liang
2018-06-23 18:32:16 -07:00
committed by GitHub
parent 9c3bab5c42
commit 0b6112b726
16 changed files with 849 additions and 277 deletions
@@ -8,14 +8,15 @@ import tensorflow as tf
import ray
from ray.rllib.models import ModelCatalog
from ray.rllib.optimizers import SampleBatch
from ray.rllib.optimizers import MultiAgentBatch
from ray.rllib.optimizers.policy_evaluator import PolicyEvaluator
from ray.rllib.utils.async_vector_env import AsyncVectorEnv, _VectorEnvToAsync
from ray.rllib.utils.async_vector_env import AsyncVectorEnv
from ray.rllib.utils.atari_wrappers import wrap_deepmind, is_atari
from ray.rllib.utils.compression import pack
from ray.rllib.utils.filter import get_filter
from ray.rllib.utils.multi_agent_env import MultiAgentEnv
from ray.rllib.utils.sampler import AsyncSampler, SyncSampler
from ray.rllib.utils.serving_env import ServingEnv, _ServingEnvToAsync
from ray.rllib.utils.serving_env import ServingEnv
from ray.rllib.utils.tf_policy_graph import TFPolicyGraph
from ray.rllib.utils.vector_env import VectorEnv
from ray.tune.result import TrainingResult
@@ -152,6 +153,7 @@ class CommonPolicyEvaluator(PolicyEvaluator):
self.env = env_creator(env_config)
if isinstance(self.env, VectorEnv) or \
isinstance(self.env, ServingEnv) or \
isinstance(self.env, MultiAgentEnv) or \
isinstance(self.env, AsyncVectorEnv):
def wrap(env):
return env # we can't auto-wrap these env types
@@ -169,8 +171,6 @@ class CommonPolicyEvaluator(PolicyEvaluator):
def make_env():
return wrap(env_creator(env_config))
self.policy_map = {}
if issubclass(policy_graph, TFPolicyGraph):
with tf.Graph().as_default():
if tf_session_creator:
@@ -186,25 +186,21 @@ class CommonPolicyEvaluator(PolicyEvaluator):
policy = policy_graph(
self.env.observation_space, self.env.action_space,
policy_config)
self.policy_map = {
"default": policy
}
self.obs_filter = get_filter(
observation_filter, self.env.observation_space.shape)
self.filters = {"obs_filter": self.obs_filter}
self.filters = {
# TODO(ekl) make the obs space dependent on policy
policy_id: get_filter(
observation_filter, self.env.observation_space.shape)
for (policy_id, policy) in self.policy_map.items()
}
# Always use vector env for consistency even if num_envs = 1
if not isinstance(self.env, AsyncVectorEnv):
if isinstance(self.env, ServingEnv):
self.vector_env = _ServingEnvToAsync(self.env)
else:
if not isinstance(self.env, VectorEnv):
self.env = VectorEnv.wrap(
make_env, [self.env], num_envs=num_envs)
self.vector_env = _VectorEnvToAsync(self.env)
else:
self.vector_env = self.env
self.async_env = AsyncVectorEnv.wrap_async(
self.env, make_env=make_env, num_envs=num_envs)
if self.batch_mode == "truncate_episodes":
if batch_steps % num_envs != 0:
@@ -222,19 +218,21 @@ class CommonPolicyEvaluator(PolicyEvaluator):
"Unsupported batch mode: {}".format(self.batch_mode))
if sample_async:
self.sampler = AsyncSampler(
self.vector_env, self.policy_map["default"], self.obs_filter,
batch_steps, horizon=episode_horizon, pack=pack_episodes)
self.async_env, self.policy_map, lambda agent_id: "default",
self.filters, batch_steps, horizon=episode_horizon,
pack=pack_episodes)
self.sampler.start()
else:
self.sampler = SyncSampler(
self.vector_env, self.policy_map["default"], self.obs_filter,
batch_steps, horizon=episode_horizon, pack=pack_episodes)
self.async_env, self.policy_map, lambda agent_id: "default",
self.filters, batch_steps, horizon=episode_horizon,
pack=pack_episodes)
def sample(self):
"""Evaluate the current policies and return a batch of experiences.
Return:
SampleBatch from evaluating the current policies.
SampleBatch|MultiAgentBatch from evaluating the current policies.
"""
batches = [self.sampler.get_data()]
@@ -243,11 +241,16 @@ class CommonPolicyEvaluator(PolicyEvaluator):
batch = self.sampler.get_data()
steps_so_far += batch.count
batches.append(batch)
batch = SampleBatch.concat_samples(batches)
batch = batches[0].concat_samples(batches)
if self.compress_observations:
batch["obs"] = [pack(o) for o in batch["obs"]]
batch["new_obs"] = [pack(o) for o in batch["new_obs"]]
if isinstance(batch, MultiAgentBatch):
for data in batch.policy_batches.values():
data["obs"] = [pack(o) for o in data["obs"]]
data["new_obs"] = [pack(o) for o in data["new_obs"]]
else:
batch["obs"] = [pack(o) for o in batch["obs"]]
batch["new_obs"] = [pack(o) for o in batch["new_obs"]]
return batch