[rllib] Part 2 of multiagent support (#2286)

* wip

* cls

* re

* wip

* wip

* a3c working

* torch support

* pg works

* lint

* rm v2

* consumer id

* clean up pg

* clean up more

* fix python 2.7

* tf session management

* docs

* dqn wip

* fix compile

* dqn

* apex runs

* up

* impotrs

* ddpg

* quotes

* fix tests

* fix last r

* fix tests

* lint

* pass checkpoint restore

* kwar

* nits

* policy graph

* fix yapf

* com

* class

* pyt

* vectorization

* update

* test cpe

* unit test

* fix ddpg2

* changes

* wip

* args

* faster test

* common

* fix

* add alg option

* batch mode and policy serving

* multi serving test

* todo

* wip

* serving test

* doc async env

* num envs

* comments

* thread

* remove init hook

* update

* fix ppo

* comments1

* fix

* updates

* add jenkins tests

* fix

* fix pytorch

* fix

* fixes

* fix a3c policy

* fix squeeze

* fix trunc on apex

* fix squeezing for real

* update

* remove horizon test for now

* multiagent wip

* update

* fix race condition

* fix ma

* t

* doc

* st

* wip

* example

* wip

* working

* cartpole

* wip

* batch wip

* fix bug

* make other_batches None default

* working

* debug

* nit

* warn

* comments

* fix ppo

* fix obs filter

* update

* fix obs filter

* pass thru worker index

* fix

* fix log action

* debug name

* fix sphinx
This commit is contained in:
Eric Liang
2018-06-25 22:33:57 -07:00
committed by GitHub
parent 739ddfa229
commit a9a26b7560
32 changed files with 939 additions and 202 deletions
+208 -47
View File
@@ -2,31 +2,38 @@ from __future__ import absolute_import
from __future__ import division
from __future__ import print_function
import pickle
import collections
import gym
import numpy as np
import pickle
import tensorflow as tf
import ray
from ray.rllib.models import ModelCatalog
from ray.rllib.optimizers import MultiAgentBatch
from ray.rllib.optimizers.policy_evaluator import PolicyEvaluator
from ray.rllib.optimizers.sample_batch import MultiAgentBatch, \
DEFAULT_POLICY_ID
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.env_context import EnvContext
from ray.rllib.utils.filter import get_filter
from ray.rllib.utils.multi_agent_env import MultiAgentEnv
from ray.rllib.utils.policy_graph import PolicyGraph
from ray.rllib.utils.sampler import AsyncSampler, SyncSampler
from ray.rllib.utils.serving_env import ServingEnv
from ray.rllib.utils.tf_policy_graph import TFPolicyGraph
from ray.rllib.utils.tf_run_builder import TFRunBuilder
from ray.rllib.utils.vector_env import VectorEnv
from ray.tune.result import TrainingResult
def collect_metrics(local_evaluator, remote_evaluators):
def collect_metrics(local_evaluator, remote_evaluators=[]):
"""Gathers episode metrics from CommonPolicyEvaluator instances."""
episode_rewards = []
episode_lengths = []
policy_rewards = collections.defaultdict(list)
metric_lists = ray.get(
[a.apply.remote(lambda ev: ev.sampler.get_metrics())
for a in remote_evaluators])
@@ -35,6 +42,8 @@ def collect_metrics(local_evaluator, remote_evaluators):
for episode in metrics:
episode_lengths.append(episode.episode_length)
episode_rewards.append(episode.episode_reward)
for (_, policy_id), reward in episode.agent_rewards.items():
policy_rewards[policy_id].append(reward)
if episode_rewards:
min_reward = min(episode_rewards)
max_reward = max(episode_rewards)
@@ -45,19 +54,22 @@ def collect_metrics(local_evaluator, remote_evaluators):
avg_length = np.mean(episode_lengths)
timesteps = np.sum(episode_lengths)
for policy_id, rewards in policy_rewards.copy().items():
policy_rewards[policy_id] = np.mean(rewards)
return TrainingResult(
episode_reward_max=max_reward,
episode_reward_min=min_reward,
episode_reward_mean=avg_reward,
episode_len_mean=avg_length,
episodes_total=len(episode_lengths),
timesteps_this_iter=timesteps)
timesteps_this_iter=timesteps,
policy_reward_mean=dict(policy_rewards))
class CommonPolicyEvaluator(PolicyEvaluator):
"""Policy evaluator implementation that operates on a rllib.PolicyGraph.
TODO: multi-agent
TODO: multi-gpu
Examples:
@@ -65,9 +77,10 @@ class CommonPolicyEvaluator(PolicyEvaluator):
>>> evaluator = CommonPolicyEvaluator(
env_creator=lambda _: gym.make("CartPole-v0"),
policy_graph=PGPolicyGraph)
>>> print(evaluator.sample().keys())
{"obs": [[...]], "actions": [[...]], "rewards": [[...]],
"dones": [[...]], "new_obs": [[...]]}
>>> print(evaluator.sample())
SampleBatch({
"obs": [[...]], "actions": [[...]], "rewards": [[...]],
"dones": [[...]], "new_obs": [[...]]})
# Creating policy evaluators using optimizer_cls.make().
>>> optimizer = LocalSyncOptimizer.make(
@@ -78,6 +91,28 @@ class CommonPolicyEvaluator(PolicyEvaluator):
},
num_workers=10)
>>> for _ in range(10): optimizer.step()
# Creating a multi-agent policy evaluator
>>> evaluator = CommonPolicyEvaluator(
env_creator=lambda _: MultiAgentTrafficGrid(num_cars=25),
policy_graph={
# Use an ensemble of two policies for car agents
"car_policy1":
(PGPolicyGraph, Box(...), Discrete(...), {"gamma": 0.99}),
"car_policy2":
(PGPolicyGraph, Box(...), Discrete(...), {"gamma": 0.95}),
# Use a single shared policy for all traffic lights
"traffic_light_policy":
(PGPolicyGraph, Box(...), Discrete(...), {}),
},
policy_mapping_fn=lambda agent_id:
random.choice(["car_policy1", "car_policy2"])
if agent_id.startswith("car_") else "traffic_light_policy")
>>> print(evaluator.sample().keys())
MultiAgentBatch({
"car_policy1": SampleBatch(...),
"car_policy2": SampleBatch(...),
"traffic_light_policy": SampleBatch(...)})
"""
@classmethod
@@ -88,6 +123,7 @@ class CommonPolicyEvaluator(PolicyEvaluator):
self,
env_creator,
policy_graph,
policy_mapping_fn=None,
tf_session_creator=None,
batch_steps=100,
batch_mode="truncate_episodes",
@@ -99,14 +135,22 @@ class CommonPolicyEvaluator(PolicyEvaluator):
observation_filter="NoFilter",
env_config=None,
model_config=None,
policy_config=None):
policy_config=None,
worker_index=0):
"""Initialize a policy evaluator.
Arguments:
env_creator (func): Function that returns a gym.Env given an
env config dict.
policy_graph (class): A class implementing rllib.PolicyGraph or
rllib.TFPolicyGraph.
EnvContext wrapped configuration.
policy_graph (class|dict): Either a class implementing
PolicyGraph, or a dictionary of policy id strings to
(PolicyGraph, obs_space, action_space, config) tuples. If a
dict is specified, then we are in multi-agent mode and a
policy_mapping_fn should also be set.
policy_mapping_fn (func): A function that maps agent ids to
policy ids in multi-agent mode. This function will be called
each time a new agent appears in an episode, to bind that agent
to a policy for the duration of the episode.
tf_session_creator (func): A function that returns a TF session.
This is optional and only useful with TFPolicyGraph.
batch_steps (int): The target number of env transitions to include
@@ -138,19 +182,26 @@ class CommonPolicyEvaluator(PolicyEvaluator):
observation_filter (str): Name of observation filter to use.
env_config (dict): Config to pass to the env creator.
model_config (dict): Config to use when creating the policy model.
policy_config (dict): Config to pass to the policy.
policy_config (dict): Config to pass to the policy. In the
multi-agent case, this config will be merged with the
per-policy configs specified by `policy_graph`.
worker_index (int): For remote evaluators, this should be set to a
non-zero and unique value. This index is passed to created envs
through EnvContext so that envs can be configured per worker.
"""
env_config = env_config or {}
env_context = EnvContext(env_config or {}, worker_index)
policy_config = policy_config or {}
model_config = model_config or {}
policy_mapping_fn = (
policy_mapping_fn or (lambda agent_id: DEFAULT_POLICY_ID))
self.env_creator = env_creator
self.policy_graph = policy_graph
self.batch_steps = batch_steps
self.batch_mode = batch_mode
self.compress_observations = compress_observations
self.env = env_creator(env_config)
self.env = env_creator(env_context)
if isinstance(self.env, VectorEnv) or \
isinstance(self.env, ServingEnv) or \
isinstance(self.env, MultiAgentEnv) or \
@@ -169,32 +220,29 @@ class CommonPolicyEvaluator(PolicyEvaluator):
self.env = wrap(self.env)
def make_env():
return wrap(env_creator(env_config))
return wrap(env_creator(env_context))
if issubclass(policy_graph, TFPolicyGraph):
self.tf_sess = None
policy_dict = _validate_and_canonicalize(policy_graph, self.env)
if _has_tensorflow_graph(policy_dict):
with tf.Graph().as_default():
if tf_session_creator:
self.sess = tf_session_creator()
self.tf_sess = tf_session_creator()
else:
self.sess = tf.Session(config=tf.ConfigProto(
self.tf_sess = tf.Session(config=tf.ConfigProto(
gpu_options=tf.GPUOptions(allow_growth=True)))
with self.sess.as_default():
policy = policy_graph(
self.env.observation_space, self.env.action_space,
policy_config)
with self.tf_sess.as_default():
self.policy_map = self._build_policy_map(
policy_dict, policy_config)
else:
policy = policy_graph(
self.env.observation_space, self.env.action_space,
policy_config)
self.policy_map = self._build_policy_map(
policy_dict, policy_config)
self.policy_map = {
"default": policy
}
self.multiagent = self.policy_map.keys() != set(DEFAULT_POLICY_ID)
self.filters = {
# TODO(ekl) make the obs space dependent on policy
policy_id: get_filter(
observation_filter, self.env.observation_space.shape)
observation_filter, policy.observation_space.shape)
for (policy_id, policy) in self.policy_map.items()
}
@@ -218,15 +266,25 @@ class CommonPolicyEvaluator(PolicyEvaluator):
"Unsupported batch mode: {}".format(self.batch_mode))
if sample_async:
self.sampler = AsyncSampler(
self.async_env, self.policy_map, lambda agent_id: "default",
self.async_env, self.policy_map, policy_mapping_fn,
self.filters, batch_steps, horizon=episode_horizon,
pack=pack_episodes)
pack=pack_episodes, tf_sess=self.tf_sess)
self.sampler.start()
else:
self.sampler = SyncSampler(
self.async_env, self.policy_map, lambda agent_id: "default",
self.async_env, self.policy_map, policy_mapping_fn,
self.filters, batch_steps, horizon=episode_horizon,
pack=pack_episodes)
pack=pack_episodes, tf_sess=self.tf_sess)
def _build_policy_map(self, policy_dict, policy_config):
policy_map = {}
for name, (cls, obs_space, act_space, conf) in sorted(
policy_dict.items()):
merged_conf = policy_config.copy()
merged_conf.update(conf)
with tf.variable_scope(name):
policy_map[name] = cls(obs_space, act_space, merged_conf)
return policy_map
def sample(self):
"""Evaluate the current policies and return a batch of experiences.
@@ -254,10 +312,15 @@ class CommonPolicyEvaluator(PolicyEvaluator):
return batch
def for_policy(self, func):
"""Apply the given function to this evaluator's default policy."""
def for_policy(self, func, policy_id=DEFAULT_POLICY_ID):
"""Apply the given function to the specified policy graph."""
return func(self.policy_map["default"])
return func(self.policy_map[policy_id])
def foreach_policy(self, func):
"""Apply the given function to each (policy, policy_id) tuple."""
return [func(policy, pid) for pid, policy in self.policy_map.items()]
def sync_filters(self, new_filters):
"""Changes self's filter to given and rebases any accumulated delta.
@@ -286,28 +349,126 @@ class CommonPolicyEvaluator(PolicyEvaluator):
return return_filters
def get_weights(self):
return self.policy_map["default"].get_weights()
return {
pid: policy.get_weights()
for pid, policy in self.policy_map.items()}
def set_weights(self, weights):
return self.policy_map["default"].set_weights(weights)
for pid, w in weights.items():
self.policy_map[pid].set_weights(w)
def compute_gradients(self, samples):
return self.policy_map["default"].compute_gradients(samples)
if isinstance(samples, MultiAgentBatch):
grad_out, info_out = {}, {}
if self.tf_sess is not None:
builder = TFRunBuilder(self.tf_sess, "compute_gradients")
for pid, batch in samples.policy_batches.items():
grad_out[pid], info_out[pid] = (
self.policy_map[pid].build_compute_gradients(
builder, batch))
grad_out = {k: builder.get(v) for k, v in grad_out.items()}
info_out = {k: builder.get(v) for k, v in info_out.items()}
else:
for pid, batch in samples.policy_batches.items():
grad_out[pid], info_out[pid] = (
self.policy_map[pid].compute_gradients(batch))
return grad_out, info_out
else:
return self.policy_map[DEFAULT_POLICY_ID].compute_gradients(
samples)
def apply_gradients(self, grads):
return self.policy_map["default"].apply_gradients(grads)
if isinstance(grads, dict):
if self.tf_sess is not None:
builder = TFRunBuilder(self.tf_sess, "apply_gradients")
outputs = {
pid: self.policy_map[pid].build_apply_gradients(
builder, grad)
for pid, grad in grads.items()
}
return {
k: builder.get(v) for k, v in outputs.items()
}
else:
return {
pid: self.policy_map[pid].apply_gradients(g)
for pid, g in grads.items()
}
else:
return self.policy_map[DEFAULT_POLICY_ID].apply_gradients(grads)
def compute_apply(self, samples):
grad_fetch, apply_fetch = self.policy_map["default"].compute_apply(
samples)
return grad_fetch
if isinstance(samples, MultiAgentBatch):
info_out = {}
if self.tf_sess is not None:
builder = TFRunBuilder(self.tf_sess, "compute_apply")
for pid, batch in samples.policy_batches.items():
info_out[pid], _ = (
self.policy_map[pid].build_compute_apply(
builder, batch))
info_out = {k: builder.get(v) for k, v in info_out.items()}
else:
for pid, batch in samples.policy_batches.items():
info_out[pid], _ = (
self.policy_map[pid].compute_apply(batch))
return info_out
else:
grad_fetch, apply_fetch = (
self.policy_map[DEFAULT_POLICY_ID].compute_apply(samples))
return grad_fetch
def save(self):
filters = self.get_filters(flush_after=True)
state = self.policy_map["default"].get_state()
state = {
pid: self.policy_map[pid].get_state()
for pid in self.policy_map
}
return pickle.dumps({"filters": filters, "state": state})
def restore(self, objs):
objs = pickle.loads(objs)
self.sync_filters(objs["filters"])
self.policy_map["default"].set_state(objs["state"])
for pid, state in objs["state"].items():
self.policy_map[pid].set_state(state)
def _validate_and_canonicalize(policy_graph, env):
if isinstance(policy_graph, dict):
for k, v in policy_graph.items():
if not isinstance(k, str):
raise ValueError(
"policy_graph keys must be strs, got {}".format(type(k)))
if not isinstance(v, tuple) or len(v) != 4:
raise ValueError(
"policy_graph values must be tuples of "
"(cls, obs_space, action_space, config), got {}".format(v))
if not issubclass(v[0], PolicyGraph):
raise ValueError(
"policy_graph tuple value 0 must be a rllib.PolicyGraph "
"class, got {}".format(v[0]))
if not isinstance(v[1], gym.Space):
raise ValueError(
"policy_graph tuple value 1 (observation_space) must be a "
"gym.Space, got {}".format(type(v[1])))
if not isinstance(v[2], gym.Space):
raise ValueError(
"policy_graph tuple value 2 (action_space) must be a "
"gym.Space, got {}".format(type(v[2])))
if not isinstance(v[3], dict):
raise ValueError(
"policy_graph tuple value 3 (config) must be a dict, "
"got {}".format(type(v[3])))
return policy_graph
elif not issubclass(policy_graph, PolicyGraph):
raise ValueError("policy_graph must be a rllib.PolicyGraph class")
else:
return {
DEFAULT_POLICY_ID: (
policy_graph, env.observation_space, env.action_space, {})}
def _has_tensorflow_graph(policy_dict):
for policy, _, _, _ in policy_dict.values():
if issubclass(policy, TFPolicyGraph):
return True
return False