mirror of
https://github.com/wassname/ray.git
synced 2026-08-11 11:24:51 +08:00
[RLlib] Support Simplex action spaces for SAC (torch and tf). (#11909)
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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) +
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user