[rllib] Support for complex / variable-length observation spaces (#8393)

This commit is contained in:
Eric Liang
2020-06-06 12:22:19 +02:00
committed by GitHub
parent b160b83d3e
commit be26a7b1b0
13 changed files with 599 additions and 15 deletions
+15
View File
@@ -220,6 +220,21 @@ Self-Supervised Model Losses
You can also use the ``custom_loss()`` API to add in self-supervised losses such as VAE reconstruction loss and L2-regularization.
Variable-length / Complex Observation Spaces
--------------------------------------------
RLlib supports complex and variable-length observation spaces, including ``gym.spaces.Tuple``, ``gym.spaces.Dict``, and ``rllib.utils.spaces.Repeated``. The handling of these spaces is transparent to the user. RLlib internally will insert preprocessors to insert padding for repeated elements, flatten complex observations into a fixed-size vector during transit, and unpack the vector into the structured tensor before sending it to the model. The flattened observation is available to the model as ``input_dict["obs_flat"]``, and the unpacked observation as ``input_dict["obs"]``.
To enable batching of struct observations, RLlib unpacks them in a `StructTensor-like format <https://github.com/tensorflow/community/blob/master/rfcs/20190910-struct-tensor.md>`__. In summary, repeated fields are "pushed down" and become the outer dimensions of tensor batches, as illustrated in this figure from the StructTensor RFC.
.. image:: struct-tensor.png
For further information about complex observation spaces, see:
* A custom environment and model that uses `repeated struct fields <https://github.com/ray-project/ray/blob/master/rllib/examples/complex_struct_space.py>`__.
* The pydoc of the `Repeated space <https://github.com/ray-project/ray/blob/master/rllib/utils/spaces/repeated.py>`__.
* The pydoc of the batched `repeated values tensor <https://github.com/ray-project/ray/blob/master/rllib/models/repeated_values.py>`__.
* The `unit tests <https://github.com/ray-project/ray/blob/master/rllib/tests/test_nested_spaces.py>`__ for Tuple and Dict spaces.
Variable-length / Parametric Action Spaces
------------------------------------------
+1
View File
@@ -79,6 +79,7 @@ Models, Preprocessors, and Action Distributions
* `Custom Action Distributions <rllib-models.html#custom-action-distributions>`__
* `Supervised Model Losses <rllib-models.html#supervised-model-losses>`__
* `Self-Supervised Model Losses <rllib-models.html#self-supervised-model-losses>`__
* `Variable-length / Complex Observation Spaces <rllib-models.html#variable-length-complex-observation-spaces>`__
* `Variable-length / Parametric Action Spaces <rllib-models.html#variable-length-parametric-action-spaces>`__
* `Autoregressive Action Distributions <rllib-models.html#autoregressive-action-distributions>`__
Binary file not shown.

After

Width:  |  Height:  |  Size: 42 KiB

+24
View File
@@ -1568,6 +1568,30 @@ py_test(
args = ["--torch", "--stop-iters=1", "--num-cpus=4"]
)
py_test(
name = "examples/complex_struct_space_tf", main = "examples/complex_struct_space.py",
tags = ["examples", "examples_C"],
size = "medium",
srcs = ["examples/complex_struct_space.py"],
args = ["--framework=tf"],
)
py_test(
name = "examples/complex_struct_space_tf_eager", main = "examples/complex_struct_space.py",
tags = ["examples", "examples_C"],
size = "medium",
srcs = ["examples/complex_struct_space.py"],
args = ["--framework=tfe"],
)
py_test(
name = "examples/complex_struct_space_torch", main = "examples/complex_struct_space.py",
tags = ["examples", "examples_C"],
size = "medium",
srcs = ["examples/complex_struct_space.py"],
args = ["--framework=torch"],
)
py_test(
name = "examples/custom_keras_model_a2c",
main = "examples/custom_keras_model.py",
+47
View File
@@ -0,0 +1,47 @@
"""Example of using variable-length Repeated / struct observation spaces.
This example shows:
- using a custom environment with Repeated / struct observations
- using a custom model to view the batched list observations
For PyTorch / TF eager mode, use the --torch and --eager flags.
"""
import argparse
import ray
from ray import tune
from ray.rllib.models import ModelCatalog
from ray.rllib.examples.env.simple_rpg import SimpleRPG
from ray.rllib.examples.models.simple_rpg_model import CustomTorchRPGModel, \
CustomTFRPGModel
parser = argparse.ArgumentParser()
parser.add_argument(
"--framework", choices=["tf", "tfe", "torch"], default="tf")
parser.add_argument("--eager", action="store_true")
if __name__ == "__main__":
ray.init()
args = parser.parse_args()
if args.framework == "torch":
ModelCatalog.register_custom_model("my_model", CustomTorchRPGModel)
else:
ModelCatalog.register_custom_model("my_model", CustomTFRPGModel)
tune.run(
"PG",
stop={
"timesteps_total": 1,
},
config={
"framework": args.framework,
"eager": args.eager,
"env": SimpleRPG,
"rollout_fragment_length": 1,
"train_batch_size": 2,
"num_workers": 0,
"model": {
"custom_model": "my_model",
},
},
)
+48
View File
@@ -0,0 +1,48 @@
import gym
from gym.spaces import Discrete, Box, Dict
from ray.rllib.utils.spaces.repeated import Repeated
# Constraints on the Repeated space.
MAX_PLAYERS = 4
MAX_ITEMS = 7
MAX_EFFECTS = 2
class SimpleRPG(gym.Env):
"""Example of a custom env with a complex, structured observation.
The observation is a list of players, each of which is a Dict of
attributes, and may further hold a list of items (categorical space).
Note that the env doesn't train, it's just a dummy example to show how to
use spaces.Repeated in a custom model (see CustomRPGModel below).
"""
def __init__(self, config):
self.cur_pos = 0
self.action_space = Discrete(4)
# Represents an item.
self.item_space = Discrete(5)
# Represents an effect on the player.
self.effect_space = Box(9000, 9999, shape=(4, ))
# Represents a player.
self.player_space = Dict({
"location": Box(-100, 100, shape=(2, )),
"status": Box(-1, 1, shape=(10, )),
"items": Repeated(self.item_space, max_len=MAX_ITEMS),
"effects": Repeated(self.effect_space, max_len=MAX_EFFECTS),
})
# Observation is a list of players.
self.observation_space = Repeated(
self.player_space, max_len=MAX_PLAYERS)
def reset(self):
return self.observation_space.sample()
def step(self, action):
return self.observation_space.sample(), 1, True, {}
+69
View File
@@ -0,0 +1,69 @@
from ray.rllib.utils import try_import_tf, try_import_torch
from ray.rllib.models.tf.tf_modelv2 import TFModelV2
from ray.rllib.models.tf.fcnet_v2 import FullyConnectedNetwork as TFFCNet
from ray.rllib.models.torch.torch_modelv2 import TorchModelV2
from ray.rllib.models.torch.fcnet import FullyConnectedNetwork as TorchFCNet
tf = try_import_tf()
torch, nn = try_import_torch()
class CustomTorchRPGModel(TorchModelV2, nn.Module):
"""Example of interpreting repeated observations."""
def __init__(self, obs_space, action_space, num_outputs, model_config,
name):
super().__init__(obs_space, action_space, num_outputs, model_config,
name)
nn.Module.__init__(self)
self.model = TorchFCNet(obs_space, action_space, num_outputs,
model_config, name)
def forward(self, input_dict, state, seq_lens):
# The unpacked input tensors, where M=MAX_PLAYERS, N=MAX_ITEMS:
# {
# 'items', <torch.Tensor shape=(?, M, N, 5)>,
# 'location', <torch.Tensor shape=(?, M, 2)>,
# 'status', <torch.Tensor shape=(?, M, 10)>,
# }
print("The unpacked input tensors:", input_dict["obs"])
print()
print("Unbatched repeat dim", input_dict["obs"].unbatch_repeat_dim())
print()
print("Fully unbatched", input_dict["obs"].unbatch_all())
print()
return self.model.forward(input_dict, state, seq_lens)
def value_function(self):
return self.model.value_function()
class CustomTFRPGModel(TFModelV2):
"""Example of interpreting repeated observations."""
def __init__(self, obs_space, action_space, num_outputs, model_config,
name):
super().__init__(obs_space, action_space, num_outputs, model_config,
name)
self.model = TFFCNet(obs_space, action_space, num_outputs,
model_config, name)
self.register_variables(self.model.variables())
def forward(self, input_dict, state, seq_lens):
# The unpacked input tensors, where M=MAX_PLAYERS, N=MAX_ITEMS:
# {
# 'items', <tf.Tensor shape=(?, M, N, 5)>,
# 'location', <tf.Tensor shape=(?, M, 2)>,
# 'status', <tf.Tensor shape=(?, M, 10)>,
# }
print("The unpacked input tensors:", input_dict["obs"])
print()
print("Unbatched repeat dim", input_dict["obs"].unbatch_repeat_dim())
print()
if tf.executing_eagerly():
print("Fully unbatched", input_dict["obs"].unbatch_all())
print()
return self.model.forward(input_dict, state, seq_lens)
def value_function(self):
return self.model.value_function()
+42 -12
View File
@@ -1,10 +1,13 @@
from collections import OrderedDict
import gym
from ray.rllib.models.preprocessors import get_preprocessor
from ray.rllib.models.preprocessors import get_preprocessor, \
RepeatedValuesPreprocessor
from ray.rllib.models.repeated_values import RepeatedValues
from ray.rllib.policy.sample_batch import SampleBatch
from ray.rllib.utils.annotations import DeveloperAPI, PublicAPI
from ray.rllib.utils.framework import try_import_tf, try_import_torch
from ray.rllib.utils.spaces.repeated import Repeated
tf = try_import_tf()
torch, _ = try_import_torch()
@@ -332,13 +335,17 @@ def _unpack_obs(obs, space, tensorlib=tf):
"""Unpack a flattened Dict or Tuple observation array/tensor.
Arguments:
obs: The flattened observation tensor
obs: The flattened observation tensor, with last dimension equal to
the flat size and any number of batch dimensions. For example, for
Box(4,), the obs may have shape [B, 4], or [B, N, M, 4] in case
the Box was nested under two Repeated spaces.
space: The original space prior to flattening
tensorlib: The library used to unflatten (reshape) the array/tensor
"""
if (isinstance(space, gym.spaces.Dict)
or isinstance(space, gym.spaces.Tuple)):
or isinstance(space, gym.spaces.Tuple)
or isinstance(space, Repeated)):
if id(space) in _cache:
prep = _cache[id(space)]
else:
@@ -346,32 +353,55 @@ def _unpack_obs(obs, space, tensorlib=tf):
# Make an attempt to cache the result, if enough space left.
if len(_cache) < 999:
_cache[id(space)] = prep
if len(obs.shape) != 2 or obs.shape[1] != prep.shape[0]:
if len(obs.shape) < 2 or obs.shape[-1] != prep.shape[0]:
raise ValueError(
"Expected flattened obs shape of [None, {}], got {}".format(
"Expected flattened obs shape of [..., {}], got {}".format(
prep.shape[0], obs.shape))
assert len(prep.preprocessors) == len(space.spaces), \
(len(prep.preprocessors) == len(space.spaces))
offset = 0
if tensorlib == tf:
batch_dims = [v.value for v in obs.shape[:-1]]
batch_dims = [-1 if v is None else v for v in batch_dims]
else:
batch_dims = list(obs.shape[:-1])
if isinstance(space, gym.spaces.Tuple):
assert len(prep.preprocessors) == len(space.spaces), \
(len(prep.preprocessors) == len(space.spaces))
u = []
for p, v in zip(prep.preprocessors, space.spaces):
obs_slice = obs[:, offset:offset + p.size]
obs_slice = obs[..., offset:offset + p.size]
offset += p.size
u.append(
_unpack_obs(
tensorlib.reshape(obs_slice, [-1] + list(p.shape)),
tensorlib.reshape(obs_slice,
batch_dims + list(p.shape)),
v,
tensorlib=tensorlib))
else:
elif isinstance(space, gym.spaces.Dict):
assert len(prep.preprocessors) == len(space.spaces), \
(len(prep.preprocessors) == len(space.spaces))
u = OrderedDict()
for p, (k, v) in zip(prep.preprocessors, space.spaces.items()):
obs_slice = obs[:, offset:offset + p.size]
obs_slice = obs[..., offset:offset + p.size]
offset += p.size
u[k] = _unpack_obs(
tensorlib.reshape(obs_slice, [-1] + list(p.shape)),
tensorlib.reshape(obs_slice, batch_dims + list(p.shape)),
v,
tensorlib=tensorlib)
elif isinstance(space, Repeated):
assert isinstance(prep, RepeatedValuesPreprocessor), prep
child_size = prep.child_preprocessor.size
# The list lengths are stored in the first slot of the flat obs.
lengths = obs[..., 0]
# [B, ..., 1 + max_len * child_sz] -> [B, ..., max_len, child_sz]
with_repeat_dim = tensorlib.reshape(
obs[..., 1:], batch_dims + [space.max_len, child_size])
# Retry the unpack, dropping the List container space.
u = _unpack_obs(
with_repeat_dim, space.child_space, tensorlib=tensorlib)
return RepeatedValues(
u, lengths=lengths, max_len=prep._obs_space.max_len)
else:
assert False, space
return u
else:
return obs
+44 -1
View File
@@ -5,6 +5,7 @@ import numpy as np
import gym
from ray.rllib.utils.annotations import override, PublicAPI
from ray.rllib.utils.spaces.repeated import Repeated
ATARI_OBS_SHAPE = (210, 160, 3)
ATARI_RAM_OBS_SHAPE = (128, )
@@ -77,7 +78,8 @@ class Preprocessor:
# Stash the unwrapped space so that we can unwrap dict and tuple spaces
# automatically in model.py
if (isinstance(self, TupleFlatteningPreprocessor)
or isinstance(self, DictFlatteningPreprocessor)):
or isinstance(self, DictFlatteningPreprocessor)
or isinstance(self, RepeatedValuesPreprocessor)):
obs_space.original_space = self._obs_space
return obs_space
@@ -243,6 +245,45 @@ class DictFlatteningPreprocessor(Preprocessor):
offset += p.size
class RepeatedValuesPreprocessor(Preprocessor):
"""Pads and batches the variable-length list value."""
@override(Preprocessor)
def _init_shape(self, obs_space, options):
assert isinstance(self._obs_space, Repeated)
child_space = obs_space.child_space
self.child_preprocessor = get_preprocessor(child_space)(child_space,
self._options)
# The first slot encodes the list length.
size = 1 + self.child_preprocessor.size * obs_space.max_len
return (size, )
@override(Preprocessor)
def transform(self, observation):
array = np.zeros(self.shape)
if isinstance(observation, list):
for elem in observation:
self.child_preprocessor.check_shape(elem)
else:
pass # ValueError will be raised in write() below.
self.write(observation, array, 0)
return array
@override(Preprocessor)
def write(self, observation, array, offset):
if not isinstance(observation, list):
raise ValueError("Input for {} must be list type, got {}".format(
self, observation))
elif len(observation) > self._obs_space.max_len:
raise ValueError("Input {} exceeds max len of space {}".format(
observation, self._obs_space.max_len))
# The first slot encodes the list length.
array[offset] = len(observation)
for i, elem in enumerate(observation):
offset_i = offset + 1 + i * self.child_preprocessor.size
self.child_preprocessor.write(elem, array, offset_i)
@PublicAPI
def get_preprocessor(space):
"""Returns an appropriate preprocessor class for the given space."""
@@ -260,6 +301,8 @@ def get_preprocessor(space):
preprocessor = TupleFlatteningPreprocessor
elif isinstance(space, gym.spaces.Dict):
preprocessor = DictFlatteningPreprocessor
elif isinstance(space, Repeated):
preprocessor = RepeatedValuesPreprocessor
else:
preprocessor = NoPreprocessor
+171
View File
@@ -0,0 +1,171 @@
from typing import List
from ray.rllib.utils.annotations import PublicAPI
from ray.rllib.utils.framework import TensorType
@PublicAPI
class RepeatedValues:
"""Represents a variable-length list of items from spaces.Repeated.
RepeatedValues are created when you use spaces.Repeated, and are
accessible as part of input_dict["obs"] in ModelV2 forward functions.
Example:
Suppose the gym space definition was:
Repeated(Repeated(Box(K), N), M)
Then in the model forward function, input_dict["obs"] is of type:
RepeatedValues(RepeatedValues(<Tensor shape=(B, M, N, K)>))
The tensor is accessible via:
input_dict["obs"].values.values
And the actual data lengths via:
# outer repetition, shape [B], range [0, M]
input_dict["obs"].lengths
-and-
# inner repetition, shape [B, M], range [0, N]
input_dict["obs"].values.lengths
Attributes:
values (Tensor): The padded data tensor of shape [B, max_len, ..., sz],
where B is the batch dimension, max_len is the max length of this
list, followed by any number of sub list max lens, followed by the
actual data size.
lengths (List[int]): Tensor of shape [B, ...] that represents the
number of valid items in each list. When the list is nested within
other lists, there will be extra dimensions for the parent list
max lens.
max_len (int): The max number of items allowed in each list.
TODO(ekl): support conversion to tf.RaggedTensor.
"""
def __init__(self, values: TensorType, lengths: List[int], max_len: int):
self.values = values
self.lengths = lengths
self.max_len = max_len
self._unbatched_repr = None
def unbatch_all(self):
"""Unbatch both the repeat and batch dimensions into Python lists.
This is only supported in PyTorch / TF eager mode.
This lets you view the data unbatched in its original form, but is
not efficient for processing.
Examples:
>>> batch = RepeatedValues(<Tensor shape=(B, N, K)>)
>>> items = batch.unbatch_all()
>>> print(len(items) == B)
True
>>> print(max(len(x) for x in items) <= N)
True
>>> print(items)
... [[<Tensor_1 shape=(K)>, ..., <Tensor_N, shape=(K)>],
... ...
... [<Tensor_1 shape=(K)>, <Tensor_2 shape=(K)>],
... ...
... [<Tensor_1 shape=(K)>],
... ...
... [<Tensor_1 shape=(K)>, ..., <Tensor_N shape=(K)>]]
"""
if self._unbatched_repr is None:
B = _get_batch_dim_helper(self.values)
if B is None:
raise ValueError(
"Cannot call unbatch_all() when batch_dim is unknown. "
"This is probably because you are using TF graph mode.")
else:
B = int(B)
slices = self.unbatch_repeat_dim()
result = []
for i in range(B):
if hasattr(self.lengths[i], "item"):
dynamic_len = int(self.lengths[i].item())
else:
dynamic_len = int(self.lengths[i].numpy())
dynamic_slice = []
for j in range(dynamic_len):
dynamic_slice.append(_batch_index_helper(slices, i, j))
result.append(dynamic_slice)
self._unbatched_repr = result
return self._unbatched_repr
def unbatch_repeat_dim(self):
"""Unbatches the repeat dimension (the one `max_len` in size).
This removes the repeat dimension. The result will be a Python list of
with length `self.max_len`. Note that the data is still padded.
Examples:
>>> batch = RepeatedValues(<Tensor shape=(B, N, K)>)
>>> items = batch.unbatch()
>>> len(items) == batch.max_len
True
>>> print(items)
... [<Tensor_1 shape=(B, K)>, ..., <Tensor_N shape=(B, K)>]
"""
return _unbatch_helper(self.values, self.max_len)
def __repr__(self):
return "RepeatedValues(value={}, lengths={}, max_len={})".format(
repr(self.values), repr(self.lengths), self.max_len)
def __str__(self):
return repr(self)
def _get_batch_dim_helper(v):
"""Tries to find the batch dimension size of v, or None."""
if isinstance(v, dict):
for u in v.values():
return _get_batch_dim_helper(u)
elif isinstance(v, tuple):
return _get_batch_dim_helper(v[0])
elif isinstance(v, RepeatedValues):
return _get_batch_dim_helper(v.values)
else:
B = v.shape[0]
if hasattr(B, "value"):
B = B.value # TensorFlow
return B
def _unbatch_helper(v, max_len):
"""Recursively unpacks the repeat dimension (max_len)."""
if isinstance(v, dict):
return {k: _unbatch_helper(u, max_len) for (k, u) in v.items()}
elif isinstance(v, tuple):
return tuple(_unbatch_helper(u, max_len) for u in v)
elif isinstance(v, RepeatedValues):
unbatched = _unbatch_helper(v.values, max_len)
return [
RepeatedValues(u, v.lengths[:, i, ...], v.max_len)
for i, u in enumerate(unbatched)
]
else:
return [v[:, i, ...] for i in range(max_len)]
def _batch_index_helper(v, i, j):
"""Selects the item at the ith batch index and jth repetition."""
if isinstance(v, dict):
return {k: _batch_index_helper(u, i, j) for (k, u) in v.items()}
elif isinstance(v, tuple):
return tuple(_batch_index_helper(u, i, j) for u in v)
elif isinstance(v, list):
# This is the output of unbatch_repeat_dim(). Unfortunately we have to
# process it here instead of in unbatch_all(), since it may be buried
# under a dict / tuple.
return _batch_index_helper(v[j], i, j)
elif isinstance(v, RepeatedValues):
unbatched = v.unbatch_all()
# Don't need to select j here; that's already done in unbatch_all.
return unbatched[i]
else:
return v[i, ...]
+97 -1
View File
@@ -1,6 +1,7 @@
from gym import spaces
from gym.envs.registration import EnvSpec
import gym
import numpy as np
import pickle
import unittest
@@ -19,6 +20,7 @@ from ray.rllib.rollout import rollout
from ray.rllib.tests.test_external_env import SimpleServing
from ray.tune.registry import register_env
from ray.rllib.utils import try_import_tf, try_import_torch
from ray.rllib.utils.spaces.repeated import Repeated
tf = try_import_tf()
_, nn = try_import_torch()
@@ -49,9 +51,23 @@ TUPLE_SPACE = spaces.Tuple([
spaces.Box(low=0, high=1, shape=(10, 10, 3)))),
spaces.Discrete(5),
])
TUPLE_SAMPLES = [TUPLE_SPACE.sample() for _ in range(10)]
# Constraints on the Repeated space.
MAX_PLAYERS = 4
MAX_ITEMS = 7
MAX_EFFECTS = 2
ITEM_SPACE = spaces.Box(-5, 5, shape=(1, ))
EFFECT_SPACE = spaces.Box(9000, 9999, shape=(4, ))
PLAYER_SPACE = spaces.Dict({
"location": spaces.Box(-100, 100, shape=(2, )),
"items": Repeated(ITEM_SPACE, max_len=MAX_ITEMS),
"effects": Repeated(EFFECT_SPACE, max_len=MAX_EFFECTS),
"status": spaces.Box(-1, 1, shape=(10, )),
})
REPEATED_SPACE = Repeated(PLAYER_SPACE, max_len=MAX_PLAYERS)
REPEATED_SAMPLES = [REPEATED_SPACE.sample() for _ in range(10)]
def one_hot(i, n):
out = [0.0] * n
@@ -91,6 +107,22 @@ class NestedTupleEnv(gym.Env):
return TUPLE_SAMPLES[self.steps], 1, self.steps >= 5, {}
class RepeatedSpaceEnv(gym.Env):
def __init__(self):
self.action_space = spaces.Discrete(2)
self.observation_space = REPEATED_SPACE
self._spec = EnvSpec("RepeatedSpaceEnv-v0")
self.steps = 0
def reset(self):
self.steps = 0
return REPEATED_SAMPLES[0]
def step(self, action):
self.steps += 1
return REPEATED_SAMPLES[self.steps], 1, self.steps >= 5, {}
class NestedMultiAgentEnv(MultiAgentEnv):
def __init__(self):
self.steps = 0
@@ -158,6 +190,45 @@ class TorchSpyModel(TorchModelV2, nn.Module):
return self.fc.value_function()
class TorchRepeatedSpyModel(TorchModelV2, nn.Module):
capture_index = 0
def __init__(self, obs_space, action_space, num_outputs, model_config,
name):
TorchModelV2.__init__(self, obs_space, action_space, num_outputs,
model_config, name)
nn.Module.__init__(self)
self.fc = FullyConnectedNetwork(
obs_space.original_space.child_space["location"], action_space,
num_outputs, model_config, name)
def forward(self, input_dict, state, seq_lens):
ray.experimental.internal_kv._internal_kv_put(
"torch_rspy_in_{}".format(TorchRepeatedSpyModel.capture_index),
pickle.dumps(input_dict["obs"].unbatch_all()),
overwrite=True)
TorchRepeatedSpyModel.capture_index += 1
return self.fc({
"obs": input_dict["obs"].values["location"][:, 0]
}, state, seq_lens)
def value_function(self):
return self.fc.value_function()
def to_list(value):
if isinstance(value, list):
return [to_list(x) for x in value]
elif isinstance(value, dict):
return {k: to_list(v) for k, v in value.items()}
elif isinstance(value, np.ndarray):
return value.tolist()
elif isinstance(value, int):
return value
else:
return value.numpy().tolist()
class DictSpyModel(TFModelV2):
capture_index = 0
@@ -435,6 +506,31 @@ class NestedSpacesTest(unittest.TestCase):
self.assertEqual(seen[1][0].tolist(), cam_i)
self.assertEqual(seen[2][0].tolist(), task_i)
# TODO(ekl) should probably also add a test for TF/eager
def test_torch_repeated(self):
ModelCatalog.register_custom_model("r1", TorchRepeatedSpyModel)
register_env("repeat", lambda _: RepeatedSpaceEnv())
a2c = A2CTrainer(
env="repeat",
config={
"num_workers": 0,
"rollout_fragment_length": 5,
"train_batch_size": 5,
"model": {
"custom_model": "r1",
},
"framework": "torch",
})
a2c.train()
# Check that the model sees the correct reconstructed observations
for i in range(4):
seen = pickle.loads(
ray.experimental.internal_kv._internal_kv_get(
"torch_rspy_in_{}".format(i)))
self.assertEqual(to_list(seen), [to_list(REPEATED_SAMPLES[i])])
if __name__ == "__main__":
import pytest
+37
View File
@@ -0,0 +1,37 @@
import numpy as np
import gym
from ray.rllib.utils.annotations import PublicAPI
@PublicAPI
class Repeated(gym.Space):
"""Represents a variable-length list of child spaces.
Example:
self.observation_space = spaces.Repeated(spaces.Box(4,), max_len=10)
--> from 0 to 10 boxes of shape (4,)
See also: documentation for rllib.models.RepeatedValues, which shows how
the lists are represented as batched input for ModelV2 classes.
"""
def __init__(self, child_space: gym.Space, max_len: int):
self.np_random = np.random.RandomState()
self.child_space = child_space
self.max_len = max_len
super().__init__()
def seed(self, seed=None):
self.np_random = np.random.RandomState()
self.np_random.seed(seed)
def sample(self):
return [
self.child_space.sample()
for _ in range(self.np_random.randint(1, self.max_len + 1))
]
def contains(self, x):
return (isinstance(x, list) and len(x) <= self.max_len
and all(self.child_space.contains(c) for c in x))
+4 -1
View File
@@ -1,7 +1,10 @@
import numpy as np
import gym
from ray.rllib.utils.annotations import PublicAPI
@PublicAPI
class Simplex(gym.Space):
"""Represents a d - 1 dimensional Simplex in R^d.
@@ -32,7 +35,7 @@ class Simplex(gym.Space):
super().__init__(shape, dtype)
self.np_random = np.random.RandomState()
def seed(self, seed):
def seed(self, seed=None):
self.np_random.seed(seed)
def sample(self):