[RLlib] Support Simplex action spaces for SAC (torch and tf). (#11909)

This commit is contained in:
Sven Mika
2020-11-11 18:45:28 +01:00
committed by GitHub
parent 4735c032ed
commit 291c172d83
8 changed files with 118 additions and 23 deletions
+9 -2
View File
@@ -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")
+9 -4
View File
@@ -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) +
+16 -8
View File
@@ -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),
+5 -1
View File
@@ -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 \
+14 -7
View File
@@ -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]
+5 -1
View File
@@ -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
+49
View File
@@ -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)
+11
View File
@@ -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)