diff --git a/rllib/BUILD b/rllib/BUILD index 1e4752be9..a92b0e9e2 100644 --- a/rllib/BUILD +++ b/rllib/BUILD @@ -1073,7 +1073,7 @@ py_test( name = "test_rollout", main = "tests/test_rollout.py", tags = ["tests_dir", "tests_dir_R"], - size = "large", + size = "enormous", data = ["train.py", "rollout.py"], srcs = ["tests/test_rollout.py"] ) diff --git a/rllib/agents/sac/sac_model.py b/rllib/agents/sac/sac_model.py index f67341fcf..d5dd0676d 100644 --- a/rllib/agents/sac/sac_model.py +++ b/rllib/agents/sac/sac_model.py @@ -161,6 +161,21 @@ class SACModel(TFModelV2): else: return self.twin_q_net(model_out) + def get_policy_output(self, model_out): + """Return the action output for the most recent forward pass. + + This outputs the support for pi(s). For continuous action spaces, this + is the action directly. For discrete, is is the mean / std dev. + + Arguments: + model_out (Tensor): obs embeddings from the model layers, of shape + [BATCH_SIZE, num_outputs]. + + Returns: + tensor of shape [BATCH_SIZE, action_out_size] + """ + return self.action_model(model_out) + def policy_variables(self): """Return the list of variables for the policy net.""" diff --git a/rllib/agents/sac/sac_policy.py b/rllib/agents/sac/sac_policy.py index 468c5ae27..09a6a5519 100644 --- a/rllib/agents/sac/sac_policy.py +++ b/rllib/agents/sac/sac_policy.py @@ -12,8 +12,8 @@ from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.policy.tf_policy import TFPolicy from ray.rllib.policy.tf_policy_template import build_tf_policy from ray.rllib.models import ModelCatalog -from ray.rllib.models.tf.tf_action_dist import ( - Categorical, SquashedGaussian, DiagGaussian) +from ray.rllib.models.tf.tf_action_dist import (Categorical, SquashedGaussian, + DiagGaussian) from ray.rllib.utils import try_import_tf, try_import_tfp from ray.rllib.utils.annotations import override from ray.rllib.utils.error import UnsupportedSpaceException @@ -95,9 +95,8 @@ def get_dist_class(config, action_space): if isinstance(action_space, Discrete): action_dist_class = Categorical else: - action_dist_class = ( - SquashedGaussian if config["normalize_actions"] - else DiagGaussian) + action_dist_class = (SquashedGaussian + if config["normalize_actions"] else DiagGaussian) return action_dist_class @@ -107,7 +106,7 @@ def get_log_likelihood(policy, model, actions, input_dict, obs_space, "obs": input_dict[SampleBatch.CUR_OBS], "is_training": policy._get_is_training_placeholder(), }, [], None) - distribution_inputs = model.action_model(model_out) + distribution_inputs = model.get_policy_output(model_out) action_dist_class = get_dist_class(policy.config, action_space) return action_dist_class(distribution_inputs, model).logp(actions) @@ -118,7 +117,7 @@ def build_action_output(policy, model, input_dict, obs_space, action_space, "obs": input_dict[SampleBatch.CUR_OBS], "is_training": policy._get_is_training_placeholder(), }, [], None) - distribution_inputs = model.action_model(model_out) + distribution_inputs = model.get_policy_output(model_out) action_dist_class = get_dist_class(policy.config, action_space) policy.output_actions, policy.sampled_action_logp = \ @@ -147,9 +146,10 @@ def actor_critic_loss(policy, model, _, train_batch): # Discrete case. if model.discrete: # Get all action probs directly from pi and form their logp. - log_pis_t = tf.nn.log_softmax(model.action_model(model_out_t), -1) + log_pis_t = tf.nn.log_softmax(model.get_policy_output(model_out_t), -1) policy_t = tf.exp(log_pis_t) - log_pis_tp1 = tf.nn.log_softmax(model.action_model(model_out_tp1), -1) + log_pis_tp1 = tf.nn.log_softmax( + model.get_policy_output(model_out_tp1), -1) policy_tp1 = tf.exp(log_pis_tp1) # Q-values. q_t = model.get_q_values(model_out_t) @@ -178,11 +178,11 @@ def actor_critic_loss(policy, model, _, train_batch): # Sample simgle actions from distribution. action_dist_class = get_dist_class(policy.config, policy.action_space) action_dist_t = action_dist_class( - model.action_model(model_out_t), policy.model) + model.get_policy_output(model_out_t), policy.model) policy_t = action_dist_t.sample() log_pis_t = tf.expand_dims(action_dist_t.sampled_action_logp(), -1) action_dist_tp1 = action_dist_class( - model.action_model(model_out_tp1), policy.model) + model.get_policy_output(model_out_tp1), policy.model) policy_tp1 = action_dist_tp1.sample() log_pis_tp1 = tf.expand_dims(action_dist_tp1.sampled_action_logp(), -1)