diff --git a/rllib/agents/ddpg/ddpg_tf_policy.py b/rllib/agents/ddpg/ddpg_tf_policy.py index 414910cc3..203add618 100644 --- a/rllib/agents/ddpg/ddpg_tf_policy.py +++ b/rllib/agents/ddpg/ddpg_tf_policy.py @@ -13,13 +13,15 @@ from ray.rllib.agents.dqn.dqn_tf_policy import postprocess_nstep_and_prio, \ PRIO_WEIGHTS from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.models import ModelCatalog -from ray.rllib.models.tf.tf_action_dist import Deterministic -from ray.rllib.models.torch.torch_action_dist import TorchDeterministic +from ray.rllib.models.tf.tf_action_dist import Deterministic, Dirichlet +from ray.rllib.models.torch.torch_action_dist import TorchDeterministic, \ + TorchDirichlet from ray.rllib.utils.annotations import override from ray.rllib.policy.tf_policy import TFPolicy 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 +from ray.rllib.utils.spaces.simplex import Simplex from ray.rllib.utils.tf_ops import huber_loss, make_tf_callable tf1, tf, tfv = try_import_tf() @@ -91,9 +93,13 @@ def get_distribution_inputs_and_class(policy, }, [], None) dist_inputs = model.get_policy_output(model_out) - return dist_inputs, (TorchDeterministic - if policy.config["framework"] == "torch" else - Deterministic), [] # []=state out + if isinstance(policy.action_space, Simplex): + distr_class = TorchDirichlet if policy.config["framework"] == "torch" \ + else Dirichlet + else: + distr_class = TorchDeterministic if \ + policy.config["framework"] == "torch" else Deterministic + return dist_inputs, distr_class, [] # []=state out def ddpg_actor_critic_loss(policy, model, _, train_batch): diff --git a/rllib/agents/ddpg/ddpg_torch_policy.py b/rllib/agents/ddpg/ddpg_torch_policy.py index f6c73f912..5041ae5fe 100644 --- a/rllib/agents/ddpg/ddpg_torch_policy.py +++ b/rllib/agents/ddpg/ddpg_torch_policy.py @@ -5,10 +5,12 @@ from ray.rllib.agents.ddpg.ddpg_tf_policy import build_ddpg_models, \ get_distribution_inputs_and_class, validate_spaces from ray.rllib.agents.dqn.dqn_tf_policy import postprocess_nstep_and_prio, \ PRIO_WEIGHTS -from ray.rllib.models.torch.torch_action_dist import TorchDeterministic +from ray.rllib.models.torch.torch_action_dist import TorchDeterministic, \ + TorchDirichlet from ray.rllib.policy.policy_template import build_policy_class from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.utils.framework import try_import_torch +from ray.rllib.utils.spaces.simplex import Simplex from ray.rllib.utils.torch_ops import apply_grad_clipping, huber_loss, l2_loss torch, nn = try_import_torch() @@ -24,7 +26,11 @@ def build_ddpg_models_and_action_dist(policy, obs_space, action_space, config): device = (torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")) policy.target_model = policy.target_model.to(device) - return model, TorchDeterministic + + if isinstance(action_space, Simplex): + return model, TorchDirichlet + else: + return model, TorchDeterministic def ddpg_actor_critic_loss(policy, model, _, train_batch): diff --git a/rllib/agents/ddpg/tests/test_ddpg.py b/rllib/agents/ddpg/tests/test_ddpg.py index 339f36fb5..0d5ddb8c5 100644 --- a/rllib/agents/ddpg/tests/test_ddpg.py +++ b/rllib/agents/ddpg/tests/test_ddpg.py @@ -184,15 +184,8 @@ class TestDDPG(unittest.TestCase): env = SimpleEnv batch_size = 100 - if env is SimpleEnv: - obs_size = (batch_size, 1) - actions = np.random.random(size=(batch_size, 1)) - elif env == "CartPole-v0": - obs_size = (batch_size, 4) - actions = np.random.randint(0, 2, size=(batch_size, )) - else: - obs_size = (batch_size, 3) - actions = np.random.random(size=(batch_size, 1)) + obs_size = (batch_size, 1) + actions = np.random.random(size=(batch_size, 1)) # Batch of size=n. input_ = self._get_batch_helper(obs_size, actions, batch_size) diff --git a/rllib/agents/sac/sac_torch_policy.py b/rllib/agents/sac/sac_torch_policy.py index d000e1839..60a206e91 100644 --- a/rllib/agents/sac/sac_torch_policy.py +++ b/rllib/agents/sac/sac_torch_policy.py @@ -32,6 +32,29 @@ F = nn.functional logger = logging.getLogger(__name__) +def _get_dist_class(config: TrainerConfigDict, action_space: gym.spaces.Space + ) -> Type[TorchDistributionWrapper]: + """Helper function to return a dist class based on config and action space. + + Args: + config (TrainerConfigDict): The Trainer's config dict. + action_space (gym.spaces.Space): The action space used. + + Returns: + Type[TFActionDistribution]: A TF distribution class. + """ + if isinstance(action_space, Discrete): + return TorchCategorical + elif isinstance(action_space, Simplex): + return TorchDirichlet + else: + if config["normalize_actions"]: + return TorchSquashedGaussian if \ + not config["_use_beta_distribution"] else TorchBeta + else: + return TorchDiagGaussian + + def build_sac_model_and_action_dist( policy: Policy, obs_space: gym.spaces.Space, @@ -56,29 +79,6 @@ def build_sac_model_and_action_dist( return model, action_dist_class -def _get_dist_class(config: TrainerConfigDict, action_space: gym.spaces.Space - ) -> Type[TorchDistributionWrapper]: - """Helper function to return a dist class based on config and action space. - - Args: - config (TrainerConfigDict): The Trainer's config dict. - action_space (gym.spaces.Space): The action space used. - - Returns: - Type[TFActionDistribution]: A TF distribution class. - """ - if isinstance(action_space, Discrete): - return TorchCategorical - elif isinstance(action_space, Simplex): - return TorchDirichlet - else: - if config["normalize_actions"]: - return TorchSquashedGaussian if \ - not config["_use_beta_distribution"] else TorchBeta - else: - return TorchDiagGaussian - - def action_distribution_fn( policy: Policy, model: ModelV2, diff --git a/rllib/agents/sac/tests/test_sac.py b/rllib/agents/sac/tests/test_sac.py index 1ec873709..b32beaac1 100644 --- a/rllib/agents/sac/tests/test_sac.py +++ b/rllib/agents/sac/tests/test_sac.py @@ -186,15 +186,8 @@ class TestSAC(unittest.TestCase): env = SimpleEnv batch_size = 100 - if env is SimpleEnv: - obs_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, )) - else: - obs_size = (batch_size, 3) - actions = np.random.random(size=(batch_size, 1)) + obs_size = (batch_size, 1) + actions = np.random.random(size=(batch_size, 2)) # Batch of size=n. input_ = self._get_batch_helper(obs_size, actions, batch_size)