[RLlib] DDPG: Support simplex action space. (#14011)

This commit is contained in:
Sven Mika
2021-02-10 15:10:01 +01:00
committed by GitHub
parent 1754359281
commit 37c7daa3c0
5 changed files with 46 additions and 48 deletions
+11 -5
View File
@@ -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):
+8 -2
View File
@@ -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):
+2 -9
View File
@@ -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)
+23 -23
View File
@@ -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,
+2 -9
View File
@@ -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)