[RLlib] Add support for custom MultiActionDistributions. (#11311)

This commit is contained in:
Sven Mika
2020-10-12 13:50:43 -07:00
committed by GitHub
parent 0c0f67c14d
commit 1ebcdf236f
2 changed files with 42 additions and 22 deletions
+41 -21
View File
@@ -3,7 +3,7 @@ import gym
import logging
import numpy as np
import tree
from typing import List
from typing import List, Optional, Type, Union
from ray.tune.registry import RLLIB_MODEL, RLLIB_PREPROCESSOR, \
RLLIB_ACTION_DIST, _global_registry
@@ -110,18 +110,20 @@ class ModelCatalog:
@staticmethod
@DeveloperAPI
def get_action_dist(action_space: gym.Space,
config: ModelConfigDict,
dist_type: str = None,
framework: str = "tf",
**kwargs) -> (type, int):
def get_action_dist(
action_space: gym.Space,
config: ModelConfigDict,
dist_type: Optional[Union[str, Type[ActionDistribution]]] = None,
framework: str = "tf",
**kwargs) -> (type, int):
"""Returns a distribution class and size for the given action space.
Args:
action_space (Space): Action space of the target gym env.
config (Optional[dict]): Optional model config.
dist_type (Optional[str]): Identifier of the action distribution
type (str) interpreted as a hint.
dist_type (Optional[Union[str, Type[ActionDistribution]]]):
Identifier of the action distribution (str) interpreted as a
hint or the actual ActionDistribution class to use.
framework (str): One of "tf", "tfe", or "torch".
kwargs (dict): Optional kwargs to pass on to the Distribution's
constructor.
@@ -143,6 +145,9 @@ class ModelCatalog:
"Using custom action distribution {}".format(action_dist_name))
dist_cls = _global_registry.get(RLLIB_ACTION_DIST,
action_dist_name)
dist_cls = ModelCatalog._get_multi_action_distribution(
dist_cls, action_space, {}, framework)
# Dist_type is given directly as a class.
elif type(dist_type) is type and \
issubclass(dist_type, ActionDistribution) and \
@@ -173,18 +178,10 @@ class ModelCatalog:
elif dist_type in (MultiActionDistribution,
TorchMultiActionDistribution) or \
isinstance(action_space, (gym.spaces.Tuple, gym.spaces.Dict)):
flat_action_space = flatten_space(action_space)
child_dists_and_in_lens = tree.map_structure(
lambda s: ModelCatalog.get_action_dist(
s, config, framework=framework), flat_action_space)
child_dists = [e[0] for e in child_dists_and_in_lens]
input_lens = [int(e[1]) for e in child_dists_and_in_lens]
return partial(
(TorchMultiActionDistribution
if framework == "torch" else MultiActionDistribution),
action_space=action_space,
child_distributions=child_dists,
input_lens=input_lens), int(sum(input_lens))
return ModelCatalog._get_multi_action_distribution(
(MultiActionDistribution
if framework == "tf" else TorchMultiActionDistribution),
action_space, config, framework)
# Simplex -> Dirichlet.
elif isinstance(action_space, Simplex):
if framework == "torch":
@@ -422,7 +419,8 @@ class ModelCatalog:
@staticmethod
@DeveloperAPI
def get_preprocessor(env: gym.Env, options: dict = None) -> Preprocessor:
def get_preprocessor(env: gym.Env,
options: Optional[dict] = None) -> Preprocessor:
"""Returns a suitable preprocessor for the given env.
This is a wrapper for get_preprocessor_for_space().
@@ -552,3 +550,25 @@ class ModelCatalog:
# Default Conv2D net.
else:
return VisionNet
@staticmethod
def _get_multi_action_distribution(dist_class, action_space, config,
framework):
# In case the custom distribution is a child of MultiActionDistr.
# If users want to completely ignore the suggested child
# distributions, they should simply do so in their custom class'
# constructor.
if issubclass(dist_class,
(MultiActionDistribution, TorchMultiActionDistribution)):
flat_action_space = flatten_space(action_space)
child_dists_and_in_lens = tree.map_structure(
lambda s: ModelCatalog.get_action_dist(
s, config, framework=framework), flat_action_space)
child_dists = [e[0] for e in child_dists_and_in_lens]
input_lens = [int(e[1]) for e in child_dists_and_in_lens]
return partial(
dist_class,
action_space=action_space,
child_distributions=child_dists,
input_lens=input_lens), int(sum(input_lens))
return dist_class
+1 -1
View File
@@ -61,7 +61,7 @@ class CustomActionDistribution(TFActionDistribution):
return tf.zeros(self.output_shape)
class ModelCatalogTest(unittest.TestCase):
class TestModelCatalog(unittest.TestCase):
def tearDown(self):
ray.shutdown()