[RLlib] Add more detailed Documentation on Model building API (#13261)

This commit is contained in:
Sven Mika
2021-01-09 12:38:29 +01:00
committed by GitHub
parent 67229bf350
commit 9dd9f72111
5 changed files with 372 additions and 156 deletions
+6 -6
View File
@@ -11,7 +11,7 @@ Available Algorithms - Overview
=================== ========== ======================= ================== =========== =============================================================
Algorithm Frameworks Discrete Actions Continuous Actions Multi-Agent Model Support
=================== ========== ======================= ================== =========== =============================================================
`A2C, A3C`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_, `+LSTM auto-wrapping`_, `+Transformer`_, `+autoreg`_
`A2C, A3C`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_, `+LSTM auto-wrapping`_, `+Attention`_, `+autoreg`_
`ARS`_ tf + torch **Yes** **Yes** No
`BC`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_
`ES`_ tf + torch **Yes** **Yes** No
@@ -20,12 +20,12 @@ Algorithm Frameworks Discrete Actions Continuous Actions Multi-
`Dreamer`_ torch No **Yes** No `+RNN`_
`DQN`_, `Rainbow`_ tf + torch **Yes** `+parametric`_ No **Yes**
`APEX-DQN`_ tf + torch **Yes** `+parametric`_ No **Yes**
`IMPALA`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_, `+LSTM auto-wrapping`_, `+Transformer`_, `+autoreg`_
`IMPALA`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_, `+LSTM auto-wrapping`_, `+Attention`_, `+autoreg`_
`MAML`_ tf + torch No **Yes** No
`MARWIL`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_
`MBMPO`_ torch No **Yes** No
`PG`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_, `+LSTM auto-wrapping`_, `+Transformer`_, `+autoreg`_
`PPO`_, `APPO`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_, `+LSTM auto-wrapping`_, `+Transformer`_, `+autoreg`_
`PG`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_, `+LSTM auto-wrapping`_, `+Attention`_, `+autoreg`_
`PPO`_, `APPO`_ tf + torch **Yes** `+parametric`_ **Yes** **Yes** `+RNN`_, `+LSTM auto-wrapping`_, `+Attention`_, `+autoreg`_
`SAC`_ tf + torch **Yes** **Yes** **Yes**
`SlateQ`_ torch **Yes** No No
`LinUCB`_, `LinTS`_ torch **Yes** `+parametric`_ No **Yes**
@@ -61,9 +61,9 @@ Algorithm Frameworks Discrete Actions Continuous Acti
.. _`+LSTM auto-wrapping`: rllib-models.html#built-in-models
.. _`+parametric`: rllib-models.html#variable-length-parametric-action-spaces
.. _`Rainbow`: rllib-algorithms.html#dqn
.. _`+RNN`: rllib-models.html#recurrent-models
.. _`+RNN`: rllib-models.html#rnns
.. _`TD3`: rllib-algorithms.html#ddpg
.. _`+Transformer`: rllib-models.html#attention-networks
.. _`+Attention`: rllib-models.html#attention
High-throughput architectures
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
+352 -138
View File
@@ -1,61 +1,201 @@
RLlib Models, Preprocessors, and Action Distributions
=====================================================
The following diagram provides a conceptual overview of data flow between different components in RLlib. We start with an ``Environment``, which given an action produces an observation. The observation is preprocessed by a ``Preprocessor`` and ``Filter`` (e.g. for running mean normalization) before being sent to a neural network ``Model``. The model output is in turn interpreted by an ``ActionDistribution`` to determine the next action.
The following diagram provides a conceptual overview of data flow between different components in RLlib.
We start with an ``Environment``, which - given an action - produces an observation.
The observation is preprocessed by a ``Preprocessor`` and ``Filter`` (e.g. for running mean normalization)
before being sent to a neural network ``Model``. The model output is in turn
interpreted by an ``ActionDistribution`` to determine the next action.
.. image:: rllib-components.svg
The components highlighted in green can be replaced with custom user-defined implementations, as described in the next sections. The purple components are RLlib internal, which means they can only be modified by changing the algorithm source code.
The components highlighted in green can be replaced with custom user-defined
implementations, as described in the next sections. The purple components are
RLlib internal, which means they can only be modified by changing the algorithm
source code.
Default Behaviours
------------------
Default Behaviors
-----------------
Built-in Preprocessors
~~~~~~~~~~~~~~~~~~~~~~
RLlib tries to pick one of its built-in preprocessor based on the environment's observation space.
Discrete observations are one-hot encoded, Atari observations downscaled, and Tuple and Dict observations flattened (these are unflattened and accessible via the ``input_dict`` parameter in custom models).
Note that for Atari, RLlib defaults to using the `DeepMind preprocessors <https://github.com/ray-project/ray/blob/master/rllib/env/atari_wrappers.py>`__, which are also used by the OpenAI baselines library.
RLlib tries to pick one of its built-in preprocessors based on the environment's
observation space. Thereby, the following simple rules apply:
Built-in Models
~~~~~~~~~~~~~~~
- Discrete observations are one-hot encoded, e.g. ``Discrete(3) and value=1 -> [0, 1, 0]``.
After preprocessing raw environment outputs, these preprocessed observations are then fed through a policy's model.
RLlib picks default models based on a simple heuristic: A vision network (`TF <https://github.com/ray-project/ray/blob/master/rllib/models/tf/visionnet.py>`__ or `Torch <https://github.com/ray-project/ray/blob/master/rllib/models/torch/visionnet.py>`__)
for observations that have a shape of length larger than 2 (for example, (84 x 84 x 3)),
and a fully connected network (`TF <https://github.com/ray-project/ray/blob/master/rllib/models/tf/fcnet.py>`__ or `Torch <https://github.com/ray-project/ray/blob/master/rllib/models/torch/fcnet.py>`__)
for everything else. These models can be configured via the ``model`` config key, documented in the model `catalog <https://github.com/ray-project/ray/blob/master/rllib/models/catalog.py>`__.
Note that for the vision network case, you'll probably have to configure ``conv_filters`` if your environment observations
have custom sizes, e.g., ``"model": {"dim": 42, "conv_filters": [[16, [4, 4], 2], [32, [4, 4], 2], [512, [11, 11], 1]]}`` for 42x42 observations.
Thereby, always make sure that the last Conv2D output has an output shape of `[B, 1, 1, X]` (`[B, X, 1, 1]` for Torch), where B=batch and
X=last Conv2D layer's number of filters, so that RLlib can flatten it. An informative error will be thrown if this is not the case.
- MultiDiscrete observations are "multi" one-hot encoded,
e.g. ``MultiDiscrete([3, 4]) and value=[1, 0] -> [0 1 0 1 0 0 0]``.
In addition, if you set ``"model": {"use_lstm": true}``, the model output will be further processed by an LSTM cell (`TF <https://github.com/ray-project/ray/blob/master/rllib/models/tf/recurrent_net.py>`__ or `Torch <https://github.com/ray-project/ray/blob/master/rllib/models/torch/recurrent_net.py>`__).
- Tuple and Dict observations are flattened, thereby, Discrete and MultiDiscrete
sub-spaces are handled as described above.
Also, the original dict/tuple observations are still available inside a) the Model via the input
dict's "obs" key (the flattened observations are in "obs_flat"), as well as b) the Policy
via the following line of code (e.g. put this into your loss function to access the original
observations: ``dict_or_tuple_obs = restore_original_dimensions(input_dict["obs"], self.obs_space, "tf|torch")``
More generally, RLlib supports the use of recurrent models for its policy gradient algorithms (A3C, PPO, PG, IMPALA), and RNN support is built into its policy evaluation utilities.
For custom RNN/LSTM setups, see the `Recurrent Models`_. section below.
For Atari observation spaces, RLlib defaults to using the `DeepMind preprocessors <https://github.com/ray-project/ray/blob/master/rllib/env/atari_wrappers.py>`__
(``preprocessor_pref=deepmind``). However, if the Trainer's config key ``preprocessor_pref`` is set to "rllib",
the following mappings apply for Atari-type observation spaces:
Built-in Model Parameters
~~~~~~~~~~~~~~~~~~~~~~~~~
- Images of shape ``(210, 160, 3)`` are downscaled to ``dim x dim``, where
``dim`` is a model config key (see default Model config below). Also, you can set
``grayscale=True`` for reducing the color channel to 1, or ``zero_mean=True`` for
producing -1.0 to 1.0 values (instead of 0.0 to 1.0 values by default).
The following is a list of the built-in model hyperparameters:
- Atari RAM observations (1D space of shape ``(128, )``) are zero-averaged
(values between -1.0 and 1.0).
In all other cases, no preprocessor will be used and the raw observations from the environment
will be sent directly into your model.
Default Model Config Settings
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
In the following paragraphs, we will first describe RLlib's default behavior for automatically constructing
models (if you don't setup a custom one), then dive into how you can customize your models by changing these
settings or writing your own model classes.
By default, RLlib will use the following config settings for your models.
These include options for the ``FullyConnectedNetworks`` (``fcnet_hiddens`` and ``fcnet_activation``),
``VisionNetworks`` (``conv_filters`` and ``conv_activation``), auto-RNN wrapping, auto-Attention (`GTrXL <https://arxiv.org/abs/1910.06764>`__) wrapping,
and some special options for Atari environments:
.. literalinclude:: ../../rllib/models/catalog.py
:language: python
:start-after: __sphinx_doc_begin__
:end-before: __sphinx_doc_end__
TensorFlow Models
-----------------
The dict above (or an overriding sub-set) is handed to the Trainer via the ``model`` key within
the main config dict like so:
.. code-block:: python
algo_config = {
# All model-related settings go into this sub-dict.
"model": {
# By default, the MODEL_DEFAULTS dict above will be used.
# Change individual keys in that dict by overriding them, e.g.
"fcnet_hiddens": [512, 512, 512],
"fcnet_activation": "relu",
},
# ... other Trainer config keys, e.g. "lr" ...
"lr": 0.00001,
}
Built-in Models
~~~~~~~~~~~~~~~
After preprocessing (if applicable) the raw environment outputs, the processed observations are fed through the policy's model.
In case, no custom model is specified (see further below on how to customize models), RLlib will pick a default model
based on simple heuristics:
- A vision network (`TF <https://github.com/ray-project/ray/blob/master/rllib/models/tf/visionnet.py>`__ or `Torch <https://github.com/ray-project/ray/blob/master/rllib/models/torch/visionnet.py>`__)
for observations that have a shape of length larger than 2, for example, ``(84 x 84 x 3)``.
- A fully connected network (`TF <https://github.com/ray-project/ray/blob/master/rllib/models/tf/fcnet.py>`__ or `Torch <https://github.com/ray-project/ray/blob/master/rllib/models/torch/fcnet.py>`__)
for everything else.
These default model types can further be configured via the ``model`` config key inside your Trainer config (as discussed above).
Available settings are `listed above <#default-model-config-settings>`__ and also documented in the `model catalog file <https://github.com/ray-project/ray/blob/master/rllib/models/catalog.py>`__.
Note that for the vision network case, you'll probably have to configure ``conv_filters``, if your environment observations
have custom sizes. For example, ``"model": {"dim": 42, "conv_filters": [[16, [4, 4], 2], [32, [4, 4], 2], [512, [11, 11], 1]]}`` for 42x42 observations.
Thereby, always make sure that the last Conv2D output has an output shape of ``[B, 1, 1, X]`` (``[B, X, 1, 1]`` for PyTorch), where B=batch and
X=last Conv2D layer's number of filters, so that RLlib can flatten it. An informative error will be thrown if this is not the case.
.. _auto_lstm_and_attention:
Built-in auto-LSTM, and auto-Attention Wrappers
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
In addition, if you set ``"use_lstm": True`` or ``"use_attention": True`` in your model config,
your model's output will be further processed by an LSTM cell
(`TF <https://github.com/ray-project/ray/blob/master/rllib/models/tf/recurrent_net.py>`__ or `Torch <https://github.com/ray-project/ray/blob/master/rllib/models/torch/recurrent_net.py>`__),
or an attention (`GTrXL <https://arxiv.org/abs/1910.06764>`__) network
(`TF <https://github.com/ray-project/ray/blob/master/rllib/models/tf/attention_net.py>`__ or
`Torch <https://github.com/ray-project/ray/blob/master/rllib/models/torch/attention_net.py>`__), respectively.
More generally, RLlib supports the use of recurrent/attention models for all
its policy-gradient algorithms (A3C, PPO, PG, IMPALA), and the necessary sequence processing support
is built into its policy evaluation utilities.
See above for which additional config keys to use to configure in more detail these two auto-wrappers
(e.g. you can specify the size of the LSTM layer by ``lstm_cell_size`` or the attention dim by ``attention_dim``).
For fully customized RNN/LSTM/Attention-Net setups see the `Recurrent Models <#rnns>`_ and
`Attention Networks/Transformers <#attention>`_ sections below.
.. note::
It is not possible to use both auto-wrappers (lstm and attention) at the same time. Doing so will create an error.
TFModelV2 replaces the previous ``rllib.models.Model`` class, which did not support Keras-style reuse of variables. The ``rllib.models.Model`` class (aka "ModelV1") is deprecated and should no longer be used.
Custom TF models should subclass `TFModelV2 <https://github.com/ray-project/ray/blob/master/rllib/models/tf/tf_modelv2.py>`__ to implement the ``__init__()`` and ``forward()`` methods. Forward takes in a dict of tensor inputs (the observation ``obs``, ``prev_action``, and ``prev_reward``, ``is_training``), optional RNN state,
and returns the model output of size ``num_outputs`` and the new state. You can also override extra methods of the model such as ``value_function`` to implement a custom value branch.
Additional supervised / self-supervised losses can be added via the ``custom_loss`` method:
Customizing Preprocessors and Models
------------------------------------
Custom Preprocessors and Environment Filters
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
.. warning::
Custom preprocessors are deprecated, since they sometimes conflict with the built-in preprocessors for handling complex observation spaces.
Please use `wrapper classes <https://github.com/openai/gym/tree/master/gym/wrappers>`__ around your environment instead of preprocessors.
Note that the built-in **default** Preprocessors described above will still be used and won't be deprecated.
Instead of using the deprecated custom Preprocessors, you should use ``gym.Wrappers`` to preprocess your environment's output (observations and rewards),
but also your Model's computed actions before sending them back to the environment.
For example, for manipulating your env's observations or rewards, do:
.. code-block:: python
import gym
from ray.rllib.utils.numpy import one_hot
class OneHotEnv(gym.core.ObservationWrapper):
# Override `observation` to custom process the original observation
# coming from the env.
def observation(self, observation):
# E.g. one-hotting a float obs [0.0, 5.0[.
return one_hot(observation, depth=5)
class ClipRewardEnv(gym.core.RewardWrapper):
def __init__(self, env, min_, max_):
super().__init__(env)
self.min = min_
self.max = max_
# Override `reward` to custom process the original reward coming
# from the env.
def reward(self, reward):
# E.g. simple clipping between min and max.
return np.clip(reward, self.min, self.max)
Custom Models: Implementing your own Forward Logic
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
If you would like to provide your own model logic (instead of using RLlib's built-in defaults), you
can sub-class either ``TFModelV2`` (for TensorFlow) or ``TorchModelV2`` (for PyTorch) and then
register and specify your sub-class in the config as follows:
Custom TensorFlow Models
````````````````````````
Custom TensorFlow models should subclass `TFModelV2 <https://github.com/ray-project/ray/blob/master/rllib/models/tf/tf_modelv2.py>`__ and implement the ``__init__()`` and ``forward()`` methods.
``forward()`` takes a dict of tensor inputs (mapping str to Tensor types), whose keys and values depend on
the `view requirements <rllib-sample-collection.html>`__ of the model.
Normally, this input dict contains only the current observation ``obs`` and an ``is_training`` boolean flag, as well as an optional list of RNN states.
``forward()`` should return the model output (of size ``self.num_outputs``) and - if applicable - a new list of internal
states (in case of RNNs or attention nets). You can also override extra methods of the model such as ``value_function`` to implement
a custom value branch.
Additional supervised/self-supervised losses can be added via the ``TFModelV2.custom_loss`` method:
.. autoclass:: ray.rllib.models.tf.tf_modelv2.TFModelV2
@@ -69,7 +209,7 @@ Additional supervised / self-supervised losses can be added via the ``custom_los
.. automethod:: variables
.. automethod:: trainable_variables
Once implemented, the model can then be registered and used in place of a built-in model:
Once implemented, your TF model can then be registered and used in place of a built-in default one:
.. code-block:: python
@@ -83,24 +223,37 @@ Once implemented, the model can then be registered and used in place of a built-
def forward(self, input_dict, state, seq_lens): ...
def value_function(self): ...
ModelCatalog.register_custom_model("my_model", MyModelClass)
ModelCatalog.register_custom_model("my_tf_model", MyModelClass)
ray.init()
trainer = ppo.PPOTrainer(env="CartPole-v0", config={
"model": {
"custom_model": "my_model",
"custom_model": "my_tf_model",
# Extra kwargs to be passed to your model's c'tor.
"custom_model_config": {},
},
})
See the `keras model example <https://github.com/ray-project/ray/blob/master/rllib/examples/custom_keras_model.py>`__ for a full example of a TF custom model.
You can also reference the `unit tests <https://github.com/ray-project/ray/blob/master/rllib/tests/test_nested_observation_spaces.py>`__ for Tuple and Dict spaces, which show how to access nested observation fields.
PyTorch Models
--------------
More examples and explanations on how to implement custom Tuple/Dict processing models
(also check out `this test case here <https://github.com/ray-project/ray/blob/master/rllib/tests/test_nested_observation_spaces.py>`__),
custom RNNs, custom model APIs (on top of default models) follow further below.
Custom PyTorch Models
`````````````````````
Similarly, you can create and register custom PyTorch models by subclassing
`TorchModelV2 <https://github.com/ray-project/ray/blob/master/rllib/models/torch/torch_modelv2.py>`__ and implement the ``__init__()`` and ``forward()`` methods.
``forward()`` takes a dict of tensor inputs (mapping str to PyTorch tensor types), whose keys and values depend on
the `view requirements <rllib-sample-collection.html>`__ of the model.
Usually, the dict contains only the current observation ``obs`` and an ``is_training`` boolean flag, as well as an optional list of RNN states.
``forward()`` should return the model output (of size ``self.num_outputs``) and - if applicable - a new list of internal
states (in case of RNNs or attention nets). You can also override extra methods of the model such as ``value_function`` to implement
a custom value branch.
Additional supervised/self-supervised losses can be added via the ``TorchModelV2.custom_loss`` method:
Similarly, you can create and register custom PyTorch models.
See these examples of `fully connected <https://github.com/ray-project/ray/blob/master/rllib/models/torch/fcnet.py>`__, `convolutional <https://github.com/ray-project/ray/blob/master/rllib/models/torch/visionnet.py>`__, and `recurrent <https://github.com/ray-project/ray/blob/master/rllib/models/torch/recurrent_net.py>`__ torch models.
.. autoclass:: ray.rllib.models.torch.torch_modelv2.TorchModelV2
@@ -111,8 +264,10 @@ See these examples of `fully connected <https://github.com/ray-project/ray/blob/
.. automethod:: custom_loss
.. automethod:: metrics
.. automethod:: get_initial_state
.. automethod:: variables
.. automethod:: trainable_variables
Once implemented, the model can then be registered and used in place of a built-in model:
Once implemented, your PyTorch model can then be registered and used in place of a built-in model:
.. code-block:: python
@@ -128,146 +283,205 @@ Once implemented, the model can then be registered and used in place of a built-
def forward(self, input_dict, state, seq_lens): ...
def value_function(self): ...
ModelCatalog.register_custom_model("my_model", CustomTorchModel)
ModelCatalog.register_custom_model("my_torch_model", CustomTorchModel)
ray.init()
trainer = ppo.PPOTrainer(env="CartPole-v0", config={
"framework": "torch",
"model": {
"custom_model": "my_model",
"custom_model": "my_torch_model",
# Extra kwargs to be passed to your model's c'tor.
"custom_model_config": {},
},
})
See the `torch model examples <https://github.com/ray-project/ray/blob/master/rllib/examples/models/>`__ for various examples on how to build a custom Torch model (including recurrent ones).
You can also reference the `unit tests <https://github.com/ray-project/ray/blob/master/rllib/tests/test_nested_observation_spaces.py>`__ for Tuple and Dict spaces, which show how to access nested observation fields.
See the `torch model examples <https://github.com/ray-project/ray/blob/master/rllib/examples/models/>`__ for various examples on how to build a custom
PyTorch model (including recurrent ones).
Recurrent Models
~~~~~~~~~~~~~~~~
More examples and explanations on how to implement custom Tuple/Dict processing models (also check out `this test case here <https://github.com/ray-project/ray/blob/master/rllib/tests/test_nested_observation_spaces.py>`__),
custom RNNs, custom model APIs (on top of default models) follow further below.
Instead of using the ``use_lstm: True`` option, it can be preferable to use a custom recurrent model.
This provides more control over postprocessing of the LSTM output and can also allow the use of multiple LSTM cells to process different portions of the input.
For an RNN model it is preferred to subclass ``RecurrentNetwork`` (either the TF or Torch versions) and to implement ``__init__()``, ``get_initial_state()``, and ``forward_rnn()``.
You can check out the `rnn_model.py <https://github.com/ray-project/ray/blob/master/rllib/examples/models/rnn_model.py>`__ models as examples to implement your own (either TF or Torch):
Wrapping a Custom Model (TF and PyTorch) with an LSTM- or Attention Net
```````````````````````````````````````````````````````````````````````
You can also use a custom (TF or PyTorch) model with our auto-wrappers for LSTMs (``use_lstm=True``) or Attention networks (``use_attention=True``).
For example, if you would like to wrap some non-default model logic with an LSTM, simply do:
.. literalinclude:: ../../rllib/examples/lstm_auto_wrapping.py
:language: python
:start-after: __sphinx_doc_begin__
:end-before: __sphinx_doc_end__
.. _rnns:
Implementing custom Recurrent Networks
``````````````````````````````````````
Instead of using the ``use_lstm: True`` option, it may be preferable to use a custom recurrent model.
This provides more control over postprocessing the LSTM's output and can also allow the use of multiple LSTM cells to process different portions of the input.
For an RNN model it is recommended to subclass ``RecurrentNetwork`` (either the `TF <https://github.com/ray-project/ray/blob/master/rllib/models/tf/recurrent_net.py>`__
or `PyTorch <https://github.com/ray-project/ray/blob/master/rllib/models/torch/recurrent_net.py>`__ versions) and then implement ``__init__()``,
``get_initial_state()``, and ``forward_rnn()``.
.. autoclass:: ray.rllib.models.tf.recurrent_net.RecurrentNetwork
.. automethod:: __init__
.. automethod:: forward_rnn
.. automethod:: get_initial_state
.. automethod:: forward_rnn
Attention Networks/Transformers
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
Note that the ``inputs`` arg entering ``forward_rnn`` is already a time-ranked single tensor (not an ``input_dict``!) with shape ``(B x T x ...)``.
If you further want to customize and need more direct access to the complete (non time-ranked) ``input_dict``, you can also override
your Model's ``forward`` method directly (as you would do with a non-RNN ModelV2). In that case, though, you are responsible for changing your inputs
and add the time rank to the incoming data (usually you just have to reshape).
RLlib now also has experimental built-in support for attention/transformer nets (the GTrXL model in particular).
Here is `an example script <https://github.com/ray-project/ray/blob/master/rllib/examples/attention_net.py>`__ on how to use these with some of our algorithms.
You can check out the `rnn_model.py <https://github.com/ray-project/ray/blob/master/rllib/examples/models/rnn_model.py>`__ models as examples to implement
your own (either TF or Torch).
.. _attention:
Implementing custom Attention Networks
``````````````````````````````````````
Similar to the RNN case described above, you could also implement your own attention-based networks, instead of using the
``use_attention: True`` flag in your model config.
Check out RLlib's `GTrXL (Attention Net) <https://arxiv.org/abs/1910.06764>`__ implementations
(for `TF <https://github.com/ray-project/ray/blob/master/rllib/models/tf/attention_net.py>`__ and `PyTorch <https://github.com/ray-project/ray/blob/master/rllib/models/torch/attention_net.py>`__)
to get a better idea on how to write your own models of this type. These are the models we use
as wrappers when ``use_attention=True``.
You can run `this example script <https://github.com/ray-project/ray/blob/master/rllib/examples/attention_net.py>`__ to run these nets within some of our algorithms.
`There is also a test case <https://github.com/ray-project/ray/blob/master/rllib/tests/test_attention_net_learning.py>`__, which confirms their learning capabilities in PPO and IMPALA.
Batch Normalization
~~~~~~~~~~~~~~~~~~~
```````````````````
You can use ``tf.layers.batch_normalization(x, training=input_dict["is_training"])`` to add batch norm layers to your custom model: `code example <https://github.com/ray-project/ray/blob/master/rllib/examples/batch_norm_model.py>`__. RLlib will automatically run the update ops for the batch norm layers during optimization (see `tf_policy.py <https://github.com/ray-project/ray/blob/master/rllib/policy/tf_policy.py>`__ and `multi_gpu_impl.py <https://github.com/ray-project/ray/blob/master/rllib/execution/multi_gpu_impl.py>`__ for the exact handling of these updates).
You can use ``tf.layers.batch_normalization(x, training=input_dict["is_training"])`` to add batch norm layers to your custom model
(see a `code example here <https://github.com/ray-project/ray/blob/master/rllib/examples/batch_norm_model.py>`__).
RLlib will automatically run the update ops for the batch norm layers during optimization
(see `tf_policy.py <https://github.com/ray-project/ray/blob/master/rllib/policy/tf_policy.py>`__ and
`multi_gpu_impl.py <https://github.com/ray-project/ray/blob/master/rllib/execution/multi_gpu_impl.py>`__ for the exact handling of these updates).
In case RLlib does not properly detect the update ops for your custom model, you can override the ``update_ops()`` method to return the list of ops to run for updates.
Custom Preprocessors
--------------------
.. warning::
Custom Model APIs (on Top of Default- or Custom Models)
```````````````````````````````````````````````````````
Custom preprocessors are deprecated, since they sometimes conflict with the built-in preprocessors for handling complex observation spaces.
Please use `wrapper classes <https://github.com/openai/gym/tree/master/gym/wrappers>`__ around your environment instead of preprocessors.
So far we talked about a) the default models that are built into RLlib and are being provided
automatically if you don't specify anything in your Trainer's config and b) custom Models through
which you can define any arbitrary forward passes.
Custom preprocessors should subclass the RLlib `preprocessor class <https://github.com/ray-project/ray/blob/master/rllib/models/preprocessors.py>`__ and be registered in the model catalog:
Another typical situation in which you would have to customize a model would be to
add a new API that your algorithm needs in order to learn, for example a Q-value
calculating head on top of your policy model. In order to expand a Model's API, simply
define and implement a new method (e.g. ``get_q_values()``) in your TF- or TorchModelV2 sub-class.
.. code-block:: python
You can now wrap this new API either around RLlib's default models or around
your custom (``forward()``-overriding) model classes. Here are two examples that illustrate how to do this:
import ray
import ray.rllib.agents.ppo as ppo
from ray.rllib.models import ModelCatalog
from ray.rllib.models.preprocessors import Preprocessor
**The Q-head API: Adding a dueling layer on top of a default RLlib model**.
class MyPreprocessorClass(Preprocessor):
def _init_shape(self, obs_space, options):
return new_shape # can vary depending on inputs
The following code adds a ``get_q_values()`` method to the automatically chosen
default Model (e.g. a ``FullyConnectedNetwork`` if the observation space is a 1D Box
or Discrete):
def transform(self, observation):
return ... # return the preprocessed observation
.. literalinclude:: ../../rllib/examples/models/custom_model_api.py
:language: python
:start-after: __sphinx_doc_model_api_1_begin__
:end-before: __sphinx_doc_model_api_1_end__
ModelCatalog.register_custom_preprocessor("my_prep", MyPreprocessorClass)
Now, for your algorithm that needs to have this model API to work properly (e.g. DQN),
you use this following code to construct the complete final Model using the
``ModelCatalog.get_model_v2`` factory function (`code here <https://github.com/ray-project/ray/blob/master/rllib/models/catalog.py>`__):
ray.init()
trainer = ppo.PPOTrainer(env="CartPole-v0", config={
"model": {
"custom_preprocessor": "my_prep",
# Extra kwargs to be passed to your model's c'tor.
"custom_model_config": {},
},
})
.. literalinclude:: ../../rllib/examples/custom_model_api.py
:language: python
:start-after: __sphinx_doc_model_construct_1_begin__
:end-before: __sphinx_doc_model_construct_1_end__
Custom Models on Top of Built-In Ones
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
A common use case is to construct a custom model on top of one of RLlib's built-in ones (e.g. a special output head on top of an fcnet, or an action + observation concat operation at the beginning or
after a conv2d stack).
Here is an example of how to construct a dueling layer head (for DQN) on top of an RLlib default model (either a Conv2D or an FCNet):
.. code-block:: python
class DuelingQModel(TFModelV2): # or: TorchModelV2
"""A simple, hard-coded dueling head model."""
def __init__(obs_space, action_space, num_outputs, model_config, name):
# Pass num_outputs=None into super constructor (so that no action/
# logits output layer is built).
# Alternatively, you can pass in num_outputs=[last layer size of
# config[model][fcnet_hiddens]] AND set no_last_linear=True, but
# this seems more tedious as you will have to explain users of this
# class that num_outputs is NOT the size of your Q-output layer.
super(DuelingQModel, self).__init__(
obs_space, action_space, None, model_config, name)
# Now: self.num_outputs contains the last layer's size, which
# we can use to construct the dueling head.
# Construct advantage head ...
self.A = tf.keras.layers.Dense(num_outputs)
# torch:
# self.A = SlimFC(
# in_size=self.num_outputs, out_size=num_outputs)
# ... and value head.
self.V = tf.keras.layers.Dense(1)
# torch:
# self.V = SlimFC(in_size=self.num_outputs, out_size=1)
def get_q_values(self, inputs):
# Calculate q-values following dueling logic:
v = self.V(inputs) # value
a = self.A(inputs) # advantages (per action)
advantages_mean = tf.reduce_mean(a, 1)
advantages_centered = a - tf.expand_dims(advantages_mean, 1)
return v + advantages_centered # q-values
With the model object constructed above, you can get the underlying intermediate output (before the dueling head)
by calling ``my_dueling_model`` directly (``out = my_dueling_model([input_dict])``), and then passing ``out`` into
your custom ``get_q_values`` method: ``q_values = my_dueling_model.get_q_values(out)``.
In order to construct an instance of the above model, you can still use the `catalog <https://github.com/ray-project/ray/blob/master/rllib/models/catalog.py>`__
`get_model_v2` convenience method:
**The single Q-value API for SAC**.
.. code-block:: python
Our DQN model from above takes an observation and outputs one Q-value per (discrete) action.
Continuous SAC - on the other hand - uses Models that calculate one Q-value only
for a single (**continuous**) action, given an observation and that particular action.
dueling_model = ModelCatalog.get_model_v2(
obs_space=[obs_space],
action_space=[action_space],
num_outputs=[num q-value (per action) outs],
model_config=config["model"],
framework="tf", # or: "torch"
model_interface=DuelingQModel,
name="dueling_q_model"
)
Let's take a look at how we would construct this API and wrap it around a custom model:
.. literalinclude:: ../../rllib/examples/models/custom_model_api.py
:language: python
:start-after: __sphinx_doc_model_api_2_begin__
:end-before: __sphinx_doc_model_api_2_end__
Now, for your algorithm that needs to have this model API to work properly (e.g. SAC),
you use this following code to construct the complete final Model using the
``ModelCatalog.get_model_v2`` factory function (`code here <https://github.com/ray-project/ray/blob/master/rllib/models/catalog.py>`__):
.. literalinclude:: ../../rllib/examples/custom_model_api.py
:language: python
:start-after: __sphinx_doc_model_construct_2_begin__
:end-before: __sphinx_doc_model_construct_2_end__
With the model object constructed above, you can get the underlying intermediate output (before the q-head)
by calling ``my_cont_action_q_model`` directly (``out = my_cont_action_q_model([input_dict])``), and then passing ``out``
and some action into your custom ``get_single_q_value`` method:
``q_value = my_cont_action_q_model.get_signle_q_value(out, action)``.
Now, with the model object, you can get the underlying intermediate output (before the dueling head)
by calling `dueling_model` directly (`out = dueling_model([input_dict])`), and then passing `out` into
your custom `get_q_values` method: `q_values = dueling_model.get_q_values(out)`.
More examples for Building Custom Models
````````````````````````````````````````
**A multi-input capable model for Tuple observation spaces (for PPO)**
RLlib's default preprocessor for Tuple and Dict spaces is to flatten incoming observations
into one flat **1D** array, and then pick a fully connected network (by default) to
process this flattened vector. This is usually ok, if you have only 1D Box or
Discrete/MultiDiscrete sub-spaces in your observations.
However, what if you had a complex observation space with one or more image components in
it (besides 1D Boxes and discrete spaces). You would probably want to preprocess each of the
image components using some convolutional network, and then concatenate their outputs
with the remaining non-image (flat) inputs (the 1D Box and discrete/one-hot components).
Take a look at this model example that does exactly that:
.. literalinclude:: ../../rllib/examples/models/cnn_plus_fc_concat_model.py
:language: python
:start-after: __sphinx_doc_begin__
:end-before: __sphinx_doc_end__
**Using the Trajectory View API: Passing in the last n actions (or rewards or observations) as inputs to a custom Model**
It is sometimes helpful for learning not only to look at the current observation
in order to calculate the next action, but also at the past n observations.
In other cases, you may want to provide the most recent rewards or actions to the model as well
(like our LSTM wrapper does if you specify: ``use_lstm=True`` and ``lstm_use_prev_action/reward=True``).
All this may even be useful when not working with partially observable environments (PO-MDPs)
and/or RNN/Attention models, as for example in classic Atari runs, where we usually use framestacking of
the last four observed images.
The `trajectory view API <rllib-sample-collection.html#trajectory-view-api>`__ allows your models
to specify these more complex "view requirements".
Here is a simple (non-RNN/Attention) example of a Model that takes as input
the last 3 observations (very similar to the recommended "framestacking" for
learning in Atari environments):
.. literalinclude:: ../../rllib/examples/models/trajectory_view_utilizing_models.py
:language: python
:start-after: __sphinx_doc_begin__
:end-before: __sphinx_doc_end__
A PyTorch version of the above model is also `given in the same file <https://github.com/ray-project/ray/blob/master/rllib/examples/models/trajectory_view_utilizing_models.py>`__.
Custom Action Distributions
@@ -505,4 +719,4 @@ To do this, you need both a custom model that implements the autoregressive patt
.. note::
Not all algorithms support autoregressive action distributions; see the `feature compatibility matrix <rllib-env.html#feature-compatibility-matrix>`__.
Not all algorithms support autoregressive action distributions; see the `algorithm overview table <rllib-algorithms.html#available-algorithms-overview>`__ for more information.
+4 -4
View File
@@ -25,7 +25,7 @@ if __name__ == "__main__":
# Run in eager mode for value checking and debugging.
tf1.enable_eager_execution()
# __sphinx_doc_model_construct_begin__
# __sphinx_doc_model_construct_1_begin__
my_dueling_model = ModelCatalog.get_model_v2(
obs_space=obs_space,
action_space=action_space,
@@ -40,7 +40,7 @@ if __name__ == "__main__":
if args.framework != "torch" else TorchDuelingQModel,
name="dueling_q_model",
)
# __sphinx_doc_model_construct_end__
# __sphinx_doc_model_construct_1_end__
batch_size = 10
input_ = np.array([obs_space.sample() for _ in range(batch_size)])
@@ -63,7 +63,7 @@ if __name__ == "__main__":
obs_space = Box(-1.0, 1.0, (3, ))
action_space = Box(-1.0, -1.0, (2, ))
# __sphinx_doc_model_construct_begin__
# __sphinx_doc_model_construct_2_begin__
my_cont_action_q_model = ModelCatalog.get_model_v2(
obs_space=obs_space,
action_space=action_space,
@@ -78,7 +78,7 @@ if __name__ == "__main__":
if args.framework != "torch" else TorchContActionQModel,
name="cont_action_q_model",
)
# __sphinx_doc_model_construct_end__
# __sphinx_doc_model_construct_2_end__
batch_size = 10
input_ = np.array([obs_space.sample() for _ in range(batch_size)])
+8 -6
View File
@@ -12,7 +12,7 @@ tf1, tf, tfv = try_import_tf()
torch, nn = try_import_torch()
# __sphinx_doc_model_api_tf_begin__
# __sphinx_doc_model_api_1_begin__
class DuelingQModel(TFModelV2): # or: TorchModelV2
"""A simple, hard-coded dueling head model."""
@@ -50,6 +50,9 @@ class DuelingQModel(TFModelV2): # or: TorchModelV2
return v + advantages_centered # q-values
# __sphinx_doc_model_api_1_end__
class TorchDuelingQModel(TorchModelV2):
"""A simple, hard-coded dueling head model."""
@@ -83,9 +86,6 @@ class TorchDuelingQModel(TorchModelV2):
return v + advantages_centered # q-values
# __sphinx_doc_model_api_tf_end__
class ContActionQModel(TFModelV2):
"""A simple, q-value-from-cont-action model (for e.g. SAC type algos)."""
@@ -127,7 +127,9 @@ class ContActionQModel(TFModelV2):
return q_values
# __sphinx_doc_model_api_torch_start__
# __sphinx_doc_model_api_2_begin__
class TorchContActionQModel(TorchModelV2):
"""A simple, q-value-from-cont-action model (for e.g. SAC type algos)."""
@@ -170,4 +172,4 @@ class TorchContActionQModel(TorchModelV2):
return q_values
# __sphinx_doc_model_api_torch_end__
# __sphinx_doc_model_api_2_end__
@@ -7,7 +7,7 @@ from ray.rllib.utils.framework import try_import_tf, try_import_torch
tf1, tf, tfv = try_import_tf()
torch, nn = try_import_torch()
# __sphinx_doc_model_api_begin__
# __sphinx_doc_begin__
class FrameStackingCartPoleModel(TFModelV2):
@@ -56,7 +56,7 @@ class FrameStackingCartPoleModel(TFModelV2):
return tf.squeeze(self._last_value, -1)
# __sphinx_doc_model_api_end__
# __sphinx_doc_end__
class TorchFrameStackingCartPoleModel(TorchModelV2, nn.Module):