From 67319bc887ad952d619cb5e02b843f64dc3f8874 Mon Sep 17 00:00:00 2001 From: Jaroslaw Rzepecki <31652222+flying-mojo@users.noreply.github.com> Date: Fri, 31 Jan 2020 20:57:52 +0000 Subject: [PATCH] [RLlib] Update MARWIL to use tf policy template (#6975) * update MARWIL to use tf policy template * formatting fixes --- rllib/agents/marwil/marwil.py | 7 +- rllib/agents/marwil/marwil_policy.py | 213 +++++++++++---------------- 2 files changed, 93 insertions(+), 127 deletions(-) diff --git a/rllib/agents/marwil/marwil.py b/rllib/agents/marwil/marwil.py index 4c288a972..ba96a14d8 100644 --- a/rllib/agents/marwil/marwil.py +++ b/rllib/agents/marwil/marwil.py @@ -1,6 +1,6 @@ from ray.rllib.agents.trainer import with_common_config from ray.rllib.agents.trainer_template import build_trainer -from ray.rllib.agents.marwil.marwil_policy import MARWILPolicy +from ray.rllib.agents.marwil.marwil_policy import MARWILTFPolicy from ray.rllib.optimizers import SyncBatchReplayOptimizer # yapf: disable @@ -53,7 +53,6 @@ def validate_config(config): MARWILTrainer = build_trainer( name="MARWIL", default_config=DEFAULT_CONFIG, - default_policy=MARWILPolicy, + default_policy=MARWILTFPolicy, validate_config=validate_config, - make_policy_optimizer=make_optimizer -) + make_policy_optimizer=make_optimizer) diff --git a/rllib/agents/marwil/marwil_policy.py b/rllib/agents/marwil/marwil_policy.py index f9a0c80d1..eae170e9e 100644 --- a/rllib/agents/marwil/marwil_policy.py +++ b/rllib/agents/marwil/marwil_policy.py @@ -1,31 +1,44 @@ +from __future__ import absolute_import +from __future__ import division +from __future__ import print_function + import ray -from ray.rllib.models import ModelCatalog +from ray.rllib.policy.sample_batch import SampleBatch +from ray.rllib.utils.explained_variance import explained_variance from ray.rllib.evaluation.postprocessing import compute_advantages, \ Postprocessing -from ray.rllib.policy.sample_batch import SampleBatch -from ray.rllib.evaluation.metrics import LEARNER_STATS_KEY -from ray.rllib.utils.annotations import override -from ray.rllib.policy.policy import Policy -from ray.rllib.policy.tf_policy import TFPolicy -from ray.rllib.utils.explained_variance import explained_variance +from ray.rllib.policy.tf_policy_template import build_tf_policy +from ray.rllib.utils.tf_ops import make_tf_callable from ray.rllib.utils import try_import_tf -from ray.rllib.utils.tf_ops import scope_vars tf = try_import_tf() -POLICY_SCOPE = "p_func" -VALUE_SCOPE = "v_func" + +class ValueNetworkMixin(object): + def __init__(self): + @make_tf_callable(self.get_session()) + def value(ob, prev_action, prev_reward, *state): + model_out, _ = self.model({ + SampleBatch.CUR_OBS: tf.convert_to_tensor([ob]), + SampleBatch.PREV_ACTIONS: tf.convert_to_tensor([prev_action]), + SampleBatch.PREV_REWARDS: tf.convert_to_tensor([prev_reward]), + "is_training": tf.convert_to_tensor(False), + }, [tf.convert_to_tensor([s]) for s in state], + tf.convert_to_tensor([1])) + return self.model.value_function()[0] + + self._value = value -class ValueLoss: +class ValueLoss(object): def __init__(self, state_values, cumulative_rewards): self.loss = 0.5 * tf.reduce_mean( tf.square(state_values - cumulative_rewards)) -class ReweightedImitationLoss: - def __init__(self, state_values, cumulative_rewards, logits, actions, - action_space, beta, model): +class ReweightedImitationLoss(object): + def __init__(self, state_values, cumulative_rewards, actions, action_dist, + beta): ma_adv_norm = tf.get_variable( name="moving_average_of_advantage_norm", dtype=tf.float32, @@ -44,130 +57,84 @@ class ReweightedImitationLoss: beta * tf.divide(adv, 1e-8 + tf.sqrt(ma_adv_norm))) # log\pi_\theta(a|s) - dist_class, _ = ModelCatalog.get_action_dist(action_space, {}) - action_dist = dist_class(logits, model) logprobs = action_dist.logp(actions) self.loss = -1.0 * tf.reduce_mean( tf.stop_gradient(exp_advs) * logprobs) -class MARWILPostprocessing: - """Adds the advantages field to the trajectory.""" +def postprocess_advantages(policy, + sample_batch, + other_agent_batches=None, + episode=None): + completed = sample_batch[SampleBatch.DONES][-1] - @override(Policy) - def postprocess_trajectory(self, - sample_batch, - other_agent_batches=None, - episode=None): - completed = sample_batch["dones"][-1] - if completed: - last_r = 0.0 - else: - raise NotImplementedError( - "last done mask in a batch should be True. " - "For now, we only support reading experience batches produced " - "with batch_mode='complete_episodes'.", - len(sample_batch[SampleBatch.DONES]), - sample_batch[SampleBatch.DONES][-1]) - batch = compute_advantages( - sample_batch, last_r, gamma=self.config["gamma"], use_gae=False) - return batch + if completed: + last_r = 0.0 + else: + next_state = [] + for i in range(policy.num_state_tensors()): + next_state.append([sample_batch["state_out_{}".format(i)][-1]]) + last_r = policy._value(sample_batch[SampleBatch.NEXT_OBS][-1], + sample_batch[SampleBatch.ACTIONS][-1], + sample_batch[SampleBatch.REWARDS][-1], + *next_state) + return compute_advantages( + sample_batch, last_r, policy.config["gamma"], use_gae=False) -class MARWILPolicy(MARWILPostprocessing, TFPolicy): - def __init__(self, observation_space, action_space, config): - config = dict(ray.rllib.agents.dqn.dqn.DEFAULT_CONFIG, **config) - self.config = config +class MARWILLoss(object): + def __init__(self, state_values, action_dist, actions, advantages, + vf_loss_coeff, beta): - dist_class, logit_dim = ModelCatalog.get_action_dist( - action_space, self.config["model"]) + self.v_loss = self._build_value_loss(state_values, advantages) + self.p_loss = self._build_policy_loss(state_values, advantages, + actions, action_dist, beta) - # Action inputs - self.obs_t = tf.placeholder( - tf.float32, shape=(None, ) + observation_space.shape) - prev_actions_ph = ModelCatalog.get_action_placeholder(action_space) - prev_rewards_ph = tf.placeholder( - tf.float32, [None], name="prev_reward") - - with tf.variable_scope(POLICY_SCOPE) as scope: - self.model = ModelCatalog.get_model({ - "obs": self.obs_t, - "prev_actions": prev_actions_ph, - "prev_rewards": prev_rewards_ph, - "is_training": self._get_is_training_placeholder(), - }, observation_space, action_space, logit_dim, - self.config["model"]) - logits = self.model.outputs - self.p_func_vars = scope_vars(scope.name) - - # Action outputs - action_dist = dist_class(logits, self.model) - self.output_actions = action_dist.sample() - - # Training inputs - self.act_t = ModelCatalog.get_action_placeholder(action_space) - self.cum_rew_t = tf.placeholder(tf.float32, [None], name="reward") - - # v network evaluation - with tf.variable_scope(VALUE_SCOPE) as scope: - state_values = self.model.value_function() - self.v_func_vars = scope_vars(scope.name) - self.v_loss = self._build_value_loss(state_values, self.cum_rew_t) - self.p_loss = self._build_policy_loss(state_values, self.cum_rew_t, - logits, self.act_t, action_space) - - # which kind of objective to optimize - objective = ( - self.p_loss.loss + self.config["vf_coeff"] * self.v_loss.loss) + self.total_loss = self.p_loss.loss + vf_loss_coeff * self.v_loss.loss self.explained_variance = tf.reduce_mean( - explained_variance(self.cum_rew_t, state_values)) - - # initialize TFPolicy - self.sess = tf.get_default_session() - self.loss_inputs = [ - (SampleBatch.CUR_OBS, self.obs_t), - (SampleBatch.ACTIONS, self.act_t), - (Postprocessing.ADVANTAGES, self.cum_rew_t), - ] - TFPolicy.__init__( - self, - observation_space, - action_space, - self.config, - self.sess, - obs_input=self.obs_t, - action_sampler=self.output_actions, - action_logp=action_dist.sampled_action_logp(), - loss=objective, - model=self.model, - loss_inputs=self.loss_inputs, - state_inputs=self.model.state_in, - state_outputs=self.model.state_out, - prev_action_input=prev_actions_ph, - prev_reward_input=prev_rewards_ph) - self.sess.run(tf.global_variables_initializer()) - - self.stats_fetches = { - "total_loss": objective, - "vf_explained_var": self.explained_variance, - "policy_loss": self.p_loss.loss, - "vf_loss": self.v_loss.loss - } + explained_variance(advantages, state_values)) def _build_value_loss(self, state_values, cum_rwds): return ValueLoss(state_values, cum_rwds) - def _build_policy_loss(self, state_values, cum_rwds, logits, actions, - action_space): - return ReweightedImitationLoss(state_values, cum_rwds, logits, actions, - action_space, self.config["beta"], - self.model) + def _build_policy_loss(self, state_values, cum_rwds, actions, action_dist, + beta): + return ReweightedImitationLoss(state_values, cum_rwds, actions, + action_dist, beta) - @override(TFPolicy) - def extra_compute_grad_fetches(self): - return {LEARNER_STATS_KEY: self.stats_fetches} - @override(Policy) - def get_initial_state(self): - return self.model.state_init +def marwil_loss(policy, model, dist_class, train_batch): + model_out, _ = model.from_batch(train_batch) + action_dist = dist_class(model_out, model) + state_values = model.value_function() + + policy.loss = MARWILLoss(state_values, action_dist, + train_batch[SampleBatch.ACTIONS], + train_batch[Postprocessing.ADVANTAGES], + policy.config["vf_coeff"], policy.config["beta"]) + + return policy.loss.total_loss + + +def stats(policy, train_batch): + return { + "policy_loss": policy.loss.p_loss.loss, + "vf_loss": policy.loss.v_loss.loss, + "total_loss": policy.loss.total_loss, + "vf_explained_var": policy.loss.explained_variance, + } + + +def setup_mixins(policy, obs_space, action_space, config): + ValueNetworkMixin.__init__(policy) + + +MARWILTFPolicy = build_tf_policy( + name="MARWILTFPolicy", + get_default_config=lambda: ray.rllib.agents.marwil.marwil.DEFAULT_CONFIG, + loss_fn=marwil_loss, + stats_fn=stats, + postprocess_fn=postprocess_advantages, + before_loss_init=setup_mixins, + mixins=[ValueNetworkMixin])