mirror of
https://github.com/wassname/ray.git
synced 2026-08-12 12:20:11 +08:00
[rllib] Improve accessing model state docs (#5656)
* [rllib] better model docs * fix * s
This commit is contained in:
@@ -207,11 +207,87 @@ Accessing Model State
|
||||
|
||||
Similar to accessing policy state, you may want to get a reference to the underlying neural network model being trained. For example, you may want to pre-train it separately, or otherwise update its weights outside of RLlib. This can be done by accessing the ``model`` of the policy:
|
||||
|
||||
**Example: Preprocessing observations for feeding into a model**
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
>>> import gym
|
||||
>>> env = gym.make("Pong-v0")
|
||||
|
||||
# RLlib uses preprocessors to implement transforms such as one-hot encoding
|
||||
# and flattening of tuple and dict observations.
|
||||
>>> from ray.rllib.models.preprocessors import get_preprocessor
|
||||
>>> prep = get_preprocessor(env.observation_space)(env.observation_space)
|
||||
<ray.rllib.models.preprocessors.GenericPixelPreprocessor object at 0x7fc4d049de80>
|
||||
|
||||
# Observations should be preprocessed prior to feeding into a model
|
||||
>>> env.reset().shape
|
||||
(210, 160, 3)
|
||||
>>> prep.transform(env.reset()).shape
|
||||
(84, 84, 3)
|
||||
|
||||
**Example: Querying a policy's action distribution**
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
# Get a reference to the policy
|
||||
>>> from ray.rllib.agents.ppo import PPOTrainer
|
||||
>>> trainer = PPOTrainer(env="CartPole-v0", config={"eager": True, "num_workers": 0})
|
||||
>>> policy = trainer.get_policy()
|
||||
<ray.rllib.policy.eager_tf_policy.PPOTFPolicy_eager object at 0x7fd020165470>
|
||||
|
||||
# Run a forward pass to get model output logits. Note that complex observations
|
||||
# must be preprocessed as in the above code block.
|
||||
>>> logits, _ = policy.model.from_batch({"obs": np.array([[0.1, 0.2, 0.3, 0.4]])})
|
||||
(<tf.Tensor: id=1274, shape=(1, 2), dtype=float32, numpy=...>, [])
|
||||
|
||||
# Compute action distribution given logits
|
||||
>>> policy.dist_class
|
||||
<class_object 'ray.rllib.models.tf.tf_action_dist.Categorical'>
|
||||
>>> dist = policy.dist_class(logits, policy.model)
|
||||
<ray.rllib.models.tf.tf_action_dist.Categorical object at 0x7fd02301d710>
|
||||
|
||||
# Query the distribution for samples, sample logps
|
||||
>>> dist.sample()
|
||||
<tf.Tensor: id=661, shape=(1,), dtype=int64, numpy=..>
|
||||
>>> dist.logp([1])
|
||||
<tf.Tensor: id=1298, shape=(1,), dtype=float32, numpy=...>
|
||||
|
||||
# Get the estimated values for the most recent forward pass
|
||||
>>> policy.model.value_function()
|
||||
<tf.Tensor: id=670, shape=(1,), dtype=float32, numpy=...>
|
||||
|
||||
>>> policy.model.base_model.summary()
|
||||
Model: "model"
|
||||
_____________________________________________________________________
|
||||
Layer (type) Output Shape Param # Connected to
|
||||
=====================================================================
|
||||
observations (InputLayer) [(None, 4)] 0
|
||||
_____________________________________________________________________
|
||||
fc_1 (Dense) (None, 256) 1280 observations[0][0]
|
||||
_____________________________________________________________________
|
||||
fc_value_1 (Dense) (None, 256) 1280 observations[0][0]
|
||||
_____________________________________________________________________
|
||||
fc_2 (Dense) (None, 256) 65792 fc_1[0][0]
|
||||
_____________________________________________________________________
|
||||
fc_value_2 (Dense) (None, 256) 65792 fc_value_1[0][0]
|
||||
_____________________________________________________________________
|
||||
fc_out (Dense) (None, 2) 514 fc_2[0][0]
|
||||
_____________________________________________________________________
|
||||
value_out (Dense) (None, 1) 257 fc_value_2[0][0]
|
||||
=====================================================================
|
||||
Total params: 134,915
|
||||
Trainable params: 134,915
|
||||
Non-trainable params: 0
|
||||
_____________________________________________________________________
|
||||
|
||||
**Example: Getting Q values from a DQN model**
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
# Get a reference to the model through the policy
|
||||
>>> from ray.rllib.agents.dqn import DQNTrainer
|
||||
>>> trainer = DQNTrainer(env="CartPole-v0")
|
||||
>>> trainer = DQNTrainer(env="CartPole-v0", config={"eager": True})
|
||||
>>> model = trainer.get_policy().model
|
||||
<ray.rllib.models.catalog.FullyConnectedNetwork_as_DistributionalQModel ...>
|
||||
|
||||
@@ -219,9 +295,10 @@ Similar to accessing policy state, you may want to get a reference to the underl
|
||||
>>> model.variables()
|
||||
[<tf.Variable 'default_policy/fc_1/kernel:0' shape=(4, 256) dtype=float32>, ...]
|
||||
|
||||
# Run a forward pass to get logits, can run with policy.get_session()
|
||||
>>> model.from_batch({"obs": np.array([[0.1, 0.2, 0.3, 0.4]])})
|
||||
(<tf.Tensor 'model_3/fc_out/Tanh:0' shape=(1, 256) dtype=float32>, [])
|
||||
# Run a forward pass to get base model output. Note that complex observations
|
||||
# must be preprocessed. An example of preprocessing is examples/saving_experiences.py
|
||||
>>> model_out = model.from_batch({"obs": np.array([[0.1, 0.2, 0.3, 0.4]])})
|
||||
(<tf.Tensor: id=832, shape=(1, 256), dtype=float32, numpy=...)
|
||||
|
||||
# Access the base Keras models (all default models have a base)
|
||||
>>> model.base_model.summary()
|
||||
@@ -243,6 +320,9 @@ Similar to accessing policy state, you may want to get a reference to the underl
|
||||
______________________________________________________________________________
|
||||
|
||||
# Access the Q value model (specific to DQN)
|
||||
>>> model.get_q_value_distributions(model_out)
|
||||
[<tf.Tensor: id=891, shape=(1, 2)>, <tf.Tensor: id=896, shape=(1, 2, 1)>]
|
||||
|
||||
>>> model.q_value_head.summary()
|
||||
Model: "model_1"
|
||||
_________________________________________________________________
|
||||
@@ -258,6 +338,9 @@ Similar to accessing policy state, you may want to get a reference to the underl
|
||||
_________________________________________________________________
|
||||
|
||||
# Access the state value model (specific to DQN)
|
||||
>>> model.get_state_value(model_out)
|
||||
<tf.Tensor: id=913, shape=(1, 1), dtype=float32>
|
||||
|
||||
>>> model.state_value_head.summary()
|
||||
Model: "model_2"
|
||||
_________________________________________________________________
|
||||
|
||||
Reference in New Issue
Block a user