diff --git a/rllib/agents/sac/sac_tf_model.py b/rllib/agents/sac/sac_tf_model.py index f8b2f6f1c..af6d1b539 100644 --- a/rllib/agents/sac/sac_tf_model.py +++ b/rllib/agents/sac/sac_tf_model.py @@ -1,10 +1,11 @@ import gym -from gym.spaces import Discrete +from gym.spaces import Box, Discrete import numpy as np from typing import Optional, Tuple from ray.rllib.models.tf.tf_modelv2 import TFModelV2 from ray.rllib.utils.framework import try_import_tf +from ray.rllib.utils.spaces.simplex import Simplex from ray.rllib.utils.typing import ModelConfigDict, TensorType tf1, tf, tfv = try_import_tf() @@ -64,11 +65,17 @@ class SACTFModel(TFModelV2): self.action_dim = action_space.n self.discrete = True action_outs = q_outs = self.action_dim - else: + elif isinstance(action_space, Box): self.action_dim = np.product(action_space.shape) self.discrete = False action_outs = 2 * self.action_dim q_outs = 1 + else: + assert isinstance(action_space, Simplex) + self.action_dim = np.product(action_space.shape) + self.discrete = False + action_outs = self.action_dim + q_outs = 1 self.model_out = tf.keras.layers.Input( shape=(self.num_outputs, ), name="model_out") diff --git a/rllib/agents/sac/sac_tf_policy.py b/rllib/agents/sac/sac_tf_policy.py index 7e4cece5e..808b7d6e4 100644 --- a/rllib/agents/sac/sac_tf_policy.py +++ b/rllib/agents/sac/sac_tf_policy.py @@ -19,13 +19,14 @@ from ray.rllib.evaluation.episode import MultiAgentEpisode from ray.rllib.models import ModelCatalog from ray.rllib.models.modelv2 import ModelV2 from ray.rllib.models.tf.tf_action_dist import Beta, Categorical, \ - DiagGaussian, SquashedGaussian, TFActionDistribution + DiagGaussian, Dirichlet, SquashedGaussian, TFActionDistribution from ray.rllib.policy.policy import Policy from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.policy.tf_policy_template import build_tf_policy from ray.rllib.utils.error import UnsupportedSpaceException from ray.rllib.utils.framework import get_variable, try_import_tf, \ try_import_tfp +from ray.rllib.utils.spaces.simplex import Simplex from ray.rllib.utils.typing import AgentID, LocalOptimizer, ModelGradients, \ TensorType, TrainerConfigDict @@ -154,6 +155,8 @@ def _get_dist_class(config: TrainerConfigDict, action_space: gym.spaces.Space """ if isinstance(action_space, Discrete): return Categorical + elif isinstance(action_space, Simplex): + return Dirichlet else: if config["normalize_actions"]: return SquashedGaussian if \ @@ -655,12 +658,14 @@ def validate_spaces(policy: Policy, observation_space: gym.spaces.Space, UnsupportedSpaceException: If one of the spaces is not supported. """ # Only support single Box or single Discreete spaces. - if not isinstance(action_space, (Box, Discrete)): + if not isinstance(action_space, (Box, Discrete, Simplex)): raise UnsupportedSpaceException( "Action space ({}) of {} is not supported for " - "SAC.".format(action_space, policy)) + "SAC. Must be [Box|Discrete|Simplex].".format( + action_space, policy)) # If Box, make sure it's a 1D vector space. - elif isinstance(action_space, Box) and len(action_space.shape) > 1: + elif isinstance(action_space, + (Box, Simplex)) and len(action_space.shape) > 1: raise UnsupportedSpaceException( "Action space ({}) of {} has multiple dimensions " "{}. ".format(action_space, policy, action_space.shape) + diff --git a/rllib/agents/sac/sac_torch_model.py b/rllib/agents/sac/sac_torch_model.py index bcd2588f0..9ebb8c75f 100644 --- a/rllib/agents/sac/sac_torch_model.py +++ b/rllib/agents/sac/sac_torch_model.py @@ -1,11 +1,12 @@ import gym -from gym.spaces import Discrete +from gym.spaces import Box, Discrete import numpy as np from typing import Optional, Tuple from ray.rllib.models.torch.misc import SlimFC from ray.rllib.models.torch.torch_modelv2 import TorchModelV2 from ray.rllib.utils.framework import get_activation_fn, try_import_torch +from ray.rllib.utils.spaces.simplex import Simplex from ray.rllib.utils.typing import ModelConfigDict, TensorType torch, nn = try_import_torch() @@ -66,13 +67,20 @@ class SACTorchModel(TorchModelV2, nn.Module): if isinstance(action_space, Discrete): self.action_dim = action_space.n self.discrete = True - self.action_outs = q_outs = self.action_dim - self.action_ins = None # No action inputs for the discrete case. - else: + action_outs = q_outs = self.action_dim + action_ins = None # No action inputs for the discrete case. + elif isinstance(action_space, Box): self.action_dim = np.product(action_space.shape) self.discrete = False - self.action_outs = 2 * self.action_dim - self.action_ins = self.action_dim + action_outs = 2 * self.action_dim + action_ins = self.action_dim + q_outs = 1 + else: + assert isinstance(action_space, Simplex) + self.action_dim = np.product(action_space.shape) + self.discrete = False + action_outs = self.action_dim + action_ins = self.action_dim q_outs = 1 # Build the policy network. @@ -94,7 +102,7 @@ class SACTorchModel(TorchModelV2, nn.Module): "action_out", SlimFC( ins, - self.action_outs, + action_outs, initializer=torch.nn.init.xavier_uniform_, activation_fn=None)) @@ -105,7 +113,7 @@ class SACTorchModel(TorchModelV2, nn.Module): # For continuous actions: Feed obs and actions (concatenated) # through the NN. For discrete actions, only obs. q_net = nn.Sequential() - ins = self.obs_ins + (0 if self.discrete else self.action_ins) + ins = self.obs_ins + (0 if self.discrete else action_ins) for i, n in enumerate(critic_hiddens): q_net.add_module( "{}_hidden_{}".format(name_, i), diff --git a/rllib/agents/sac/sac_torch_policy.py b/rllib/agents/sac/sac_torch_policy.py index 42a796a2e..5b39137b7 100644 --- a/rllib/agents/sac/sac_torch_policy.py +++ b/rllib/agents/sac/sac_torch_policy.py @@ -14,13 +14,15 @@ from ray.rllib.agents.sac.sac_tf_policy import build_sac_model, \ postprocess_trajectory, validate_spaces from ray.rllib.agents.dqn.dqn_tf_policy import PRIO_WEIGHTS from ray.rllib.models.modelv2 import ModelV2 -from ray.rllib.models.torch.torch_action_dist import TorchDistributionWrapper +from ray.rllib.models.torch.torch_action_dist import \ + TorchDistributionWrapper, TorchDirichlet from ray.rllib.policy.policy import Policy from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.policy.torch_policy_template import build_torch_policy from ray.rllib.models.torch.torch_action_dist import ( TorchCategorical, TorchSquashedGaussian, TorchDiagGaussian, TorchBeta) from ray.rllib.utils.framework import try_import_torch +from ray.rllib.utils.spaces.simplex import Simplex from ray.rllib.utils.typing import LocalOptimizer, TensorType, \ TrainerConfigDict @@ -67,6 +69,8 @@ def _get_dist_class(config: TrainerConfigDict, action_space: gym.spaces.Space """ if isinstance(action_space, Discrete): return TorchCategorical + elif isinstance(action_space, Simplex): + return TorchDirichlet else: if config["normalize_actions"]: return TorchSquashedGaussian if \ diff --git a/rllib/agents/sac/tests/test_sac.py b/rllib/agents/sac/tests/test_sac.py index 30055b30d..6d473179e 100644 --- a/rllib/agents/sac/tests/test_sac.py +++ b/rllib/agents/sac/tests/test_sac.py @@ -9,12 +9,13 @@ import ray.rllib.agents.sac as sac from ray.rllib.agents.sac.sac_tf_policy import sac_actor_critic_loss as tf_loss from ray.rllib.agents.sac.sac_torch_policy import actor_critic_loss as \ loss_torch -from ray.rllib.models.tf.tf_action_dist import SquashedGaussian -from ray.rllib.models.torch.torch_action_dist import TorchSquashedGaussian +from ray.rllib.models.tf.tf_action_dist import Dirichlet +from ray.rllib.models.torch.torch_action_dist import TorchDirichlet from ray.rllib.execution.replay_buffer import LocalReplayBuffer from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.utils.framework import try_import_tf, try_import_torch from ray.rllib.utils.numpy import fc, relu +from ray.rllib.utils.spaces.simplex import Simplex from ray.rllib.utils.test_utils import check, check_compute_single_action, \ framework_iterator from ray.rllib.utils.torch_ops import convert_to_torch_tensor @@ -25,7 +26,10 @@ torch, _ = try_import_torch() class SimpleEnv(Env): def __init__(self, config): - self.action_space = Box(0.0, 1.0, (1, )) + if config.get("simplex_actions", False): + self.action_space = Simplex((2, )) + else: + self.action_space = Box(0.0, 1.0, (1, )) self.observation_space = Box(0.0, 1.0, (1, )) self.max_steps = config.get("max_steps", 100) self.state = None @@ -38,8 +42,8 @@ class SimpleEnv(Env): def step(self, action): self.steps += 1 - # Reward is 1.0 - (action - state). - [r] = 1.0 - np.abs(action - self.state) + # Reward is 1.0 - (max(actions) - state). + [r] = 1.0 - np.abs(np.max(action) - self.state) d = self.steps >= self.max_steps self.state = self.observation_space.sample() return self.state, r, d, {} @@ -95,6 +99,8 @@ class TestSAC(unittest.TestCase): config["policy_model"]["fcnet_hiddens"] = [10] # Make sure, timing differences do not affect trainer.train(). config["min_iter_time_s"] = 0 + # Test SAC with Simplex action space. + config["env_config"] = {"simplex_actions": True} map_ = { # Normal net. @@ -147,7 +153,7 @@ class TestSAC(unittest.TestCase): batch_size = 100 if env is SimpleEnv: obs_size = (batch_size, 1) - actions = np.random.random(size=(batch_size, 1)) + actions = np.random.random(size=(batch_size, 2)) elif env == "CartPole-v0": obs_size = (batch_size, 4) actions = np.random.randint(0, 2, size=(batch_size, )) @@ -419,7 +425,8 @@ class TestSAC(unittest.TestCase): # 16=target Q out bias # 17=target Q out kernel alpha = np.exp(log_alpha) - cls = TorchSquashedGaussian if fw == "torch" else SquashedGaussian + # cls = TorchSquashedGaussian if fw == "torch" else SquashedGaussian + cls = TorchDirichlet if fw == "torch" else Dirichlet model_out_t = train_batch[SampleBatch.CUR_OBS] model_out_tp1 = train_batch[SampleBatch.NEXT_OBS] target_model_out_tp1 = train_batch[SampleBatch.NEXT_OBS] diff --git a/rllib/models/tf/tf_action_dist.py b/rllib/models/tf/tf_action_dist.py index 1647d2c0e..d125b99d0 100644 --- a/rllib/models/tf/tf_action_dist.py +++ b/rllib/models/tf/tf_action_dist.py @@ -513,13 +513,17 @@ class Dirichlet(TFActionDistribution): """ self.epsilon = 1e-7 concentration = tf.exp(inputs) + self.epsilon - self.dist = tf.distributions.Dirichlet( + self.dist = tf1.distributions.Dirichlet( concentration=concentration, validate_args=True, allow_nan_stats=False, ) super().__init__(concentration, model) + @override(ActionDistribution) + def deterministic_sample(self) -> TensorType: + return tf.nn.softmax(self.dist.concentration) + @override(ActionDistribution) def logp(self, x): # Support of Dirichlet are positive real numbers. x is already diff --git a/rllib/models/torch/torch_action_dist.py b/rllib/models/torch/torch_action_dist.py index bce456944..57b197b80 100644 --- a/rllib/models/torch/torch_action_dist.py +++ b/rllib/models/torch/torch_action_dist.py @@ -429,3 +429,52 @@ class TorchMultiActionDistribution(TorchDistributionWrapper): @override(ActionDistribution) def required_model_output_shape(self, action_space, model_config): return np.sum(self.input_lens) + + +class TorchDirichlet(TorchDistributionWrapper): + """Dirichlet distribution for continuous actions that are between + [0,1] and sum to 1. + + e.g. actions that represent resource allocation.""" + + def __init__(self, inputs, model): + """Input is a tensor of logits. The exponential of logits is used to + parametrize the Dirichlet distribution as all parameters need to be + positive. An arbitrary small epsilon is added to the concentration + parameters to be zero due to numerical error. + + See issue #4440 for more details. + """ + self.epsilon = torch.tensor(1e-7).to(inputs.device) + concentration = torch.exp(inputs) + self.epsilon + self.dist = torch.distributions.dirichlet.Dirichlet( + concentration=concentration, + validate_args=True, + ) + super().__init__(concentration, model) + + @override(ActionDistribution) + def deterministic_sample(self) -> TensorType: + return nn.functional.softmax(self.dist.concentration) + + @override(ActionDistribution) + def logp(self, x): + # Support of Dirichlet are positive real numbers. x is already + # an array of positive numbers, but we clip to avoid zeros due to + # numerical errors. + x = torch.max(x, self.epsilon) + x = x / torch.sum(x, dim=-1, keepdim=True) + return self.dist.log_prob(x) + + @override(ActionDistribution) + def entropy(self): + return self.dist.entropy() + + @override(ActionDistribution) + def kl(self, other): + return self.dist.kl_divergence(other.dist) + + @staticmethod + @override(ActionDistribution) + def required_model_output_shape(action_space, model_config): + return np.prod(action_space.shape) diff --git a/rllib/utils/exploration/random.py b/rllib/utils/exploration/random.py index 1a5549be9..5fcc344d5 100644 --- a/rllib/utils/exploration/random.py +++ b/rllib/utils/exploration/random.py @@ -10,6 +10,7 @@ from ray.rllib.utils.exploration.exploration import Exploration from ray.rllib.utils import force_tuple from ray.rllib.utils.framework import try_import_tf, try_import_torch, \ TensorType +from ray.rllib.utils.spaces.simplex import Simplex from ray.rllib.utils.spaces.space_utils import get_base_struct_from_space tf1, tf, tfv = try_import_tf() @@ -96,6 +97,16 @@ class Random(Exploration): return tf.random.normal( shape=(batch_size, ) + component.shape, dtype=component.dtype) + else: + assert isinstance(component, Simplex), \ + "Unsupported distribution component '{}' for random " \ + "sampling!".format(component) + return tf.nn.softmax( + tf.random.uniform( + shape=(batch_size, ) + component.shape, + minval=0.0, + maxval=1.0, + dtype=component.dtype)) actions = tree.map_structure(random_component, self.action_space_struct)