mirror of
https://github.com/wassname/ray.git
synced 2026-08-08 11:25:28 +08:00
[rllib] dqn/ddpg policy customization (#2445)
* dqn policy update - more customization * docs for custom DQN graph * Update rllib-training.rst * Update rllib-models.rst * Update rllib.rst * Update rllib-training.rst * Update rllib-concepts.rst * yapf codestyle
This commit is contained in:
committed by
Eric Liang
parent
68660453e4
commit
05490b8cb9
@@ -15,7 +15,7 @@ Most interaction with deep learning frameworks is isolated to the `PolicyGraph i
|
||||
Policy Evaluation
|
||||
-----------------
|
||||
|
||||
Given an environment and policy graph, policy evaluation produces `batches <https://github.com/ray-project/ray/blob/master/python/ray/rllib/evaluation/sample_batch.py>`__ of experiences. This is your classic "environment interaction loop". Efficient policy evaluation can be burdensome to get right, especially when leveraging vectorization, RNNs, or when operating in a multi-agent environment. RLlib provides a `PolicyEvaluator <https://github.com/ray-project/ray/blob/master/python/ray/rllib/evaluation/policy_evaluator.py>`__ class that manages all of this, and this class is used in most RLlib algorithm.
|
||||
Given an environment and policy graph, policy evaluation produces `batches <https://github.com/ray-project/ray/blob/master/python/ray/rllib/evaluation/sample_batch.py>`__ of experiences. This is your classic "environment interaction loop". Efficient policy evaluation can be burdensome to get right, especially when leveraging vectorization, RNNs, or when operating in a multi-agent environment. RLlib provides a `PolicyEvaluator <https://github.com/ray-project/ray/blob/master/python/ray/rllib/evaluation/policy_evaluator.py>`__ class that manages all of this, and this class is used in most RLlib algorithms.
|
||||
|
||||
You can also use policy evaluation standalone to produce batches of experiences. This can be done by calling ``ev.sample()`` on an evaluator instance, or ``ev.sample.remote()`` in parallel on evaluator instances created as Ray actors (see ``PolicyEvalutor.as_remote()``).
|
||||
|
||||
|
||||
@@ -75,3 +75,53 @@ Similarly, custom preprocessors should subclass the RLlib `preprocessor class <h
|
||||
"custom_options": {}, # extra options to pass to your preprocessor
|
||||
},
|
||||
})
|
||||
|
||||
|
||||
Customizing Policy Graphs
|
||||
-------------------------
|
||||
|
||||
For deeper customization of algorithms, you can modify the policy graphs of the agent classes. Here's an example of extending the DDPG policy graph to specify custom sub-network modules:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from ray.rllib.models import ModelCatalog
|
||||
from ray.rllib.agents.ddpg.ddpg_policy_graph import DDPGPolicyGraph as BaseDDPGPolicyGraph
|
||||
|
||||
class CustomPNetwork(object):
|
||||
def __init__(self, dim_actions, hiddens, activation):
|
||||
action_out = ...
|
||||
# Use sigmoid layer to bound values within (0, 1)
|
||||
# shape of action_scores is [batch_size, dim_actions]
|
||||
self.action_scores = layers.fully_connected(
|
||||
action_out, num_outputs=dim_actions, activation_fn=tf.nn.sigmoid)
|
||||
|
||||
class CustomQNetwork(object):
|
||||
def __init__(self, action_inputs, hiddens, activation):
|
||||
q_out = ...
|
||||
self.value = layers.fully_connected(
|
||||
q_out, num_outputs=1, activation_fn=None)
|
||||
|
||||
class CustomDDPGPolicyGraph(BaseDDPGPolicyGraph):
|
||||
def _build_p_network(self, obs):
|
||||
return CustomPNetwork(
|
||||
self.dim_actions,
|
||||
self.config["actor_hiddens"],
|
||||
self.config["actor_hidden_activation"]).action_scores
|
||||
|
||||
def _build_q_network(self, obs, actions):
|
||||
return CustomQNetwork(
|
||||
actions,
|
||||
self.config["critic_hiddens"],
|
||||
self.config["critic_hidden_activation"]).value
|
||||
|
||||
Then, you can create an agent with your custom policy graph by:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from ray.rllib.agents.ddpg.ddpg import DDPGAgent
|
||||
from custom_policy_graph import CustomDDPGPolicyGraph
|
||||
|
||||
DDPGAgent._policy_graph = CustomDDPGPolicyGraph
|
||||
agent = DDPGAgent(...)
|
||||
|
||||
That's it. In this example we overrode existing methods of the existing DDPG policy graph, i.e., `_build_q_network`, `_build_p_network`, `_build_action_network`, `_build_actor_critic_loss`, but you can also replace the entire graph class entirely.
|
||||
|
||||
@@ -56,6 +56,7 @@ Models and Preprocessors
|
||||
* `Built-in Models and Preprocessors <rllib-models.html#built-in-models-and-preprocessors>`__
|
||||
* `Custom Models <rllib-models.html#custom-models>`__
|
||||
* `Custom Preprocessors <rllib-models.html#custom-preprocessors>`__
|
||||
* `Customizing Policy Graphs <rllib-models.html#customizing-policy-graphs>`__
|
||||
|
||||
RLlib Concepts
|
||||
--------------
|
||||
|
||||
Reference in New Issue
Block a user