mirror of
https://github.com/wassname/ray.git
synced 2026-08-01 12:51:09 +08:00
[RLlib] DDPG: Support simplex action space. (#14011)
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user