mirror of
https://github.com/wassname/ray.git
synced 2026-10-01 12:31:37 +08:00
[rllib] Add Decentralized DDPPO trainer and documentation (#7088)
This commit is contained in:
1 parent
6e1c3ea824
commit
026f6884b5
15 files changed
+212
-72
No files matched your search
@@ -0,0 +1 @@
|
||||
<svg version="1.1" viewBox="0.0 0.0 631.0 144.0" fill="none" stroke="none" stroke-linecap="square" stroke-miterlimit="10" xmlns:xlink="http://www.w3.org/1999/xlink" xmlns="http://www.w3.org/2000/svg"><clipPath id="p.0"><path d="m0 0l631.0 0l0 144.0l-631.0 0l0 -144.0z" clip-rule="nonzero"/></clipPath><g clip-path="url(#p.0)"><path fill="#000000" fill-opacity="0.0" d="m0 0l631.0 0l0 144.0l-631.0 0z" fill-rule="evenodd"/><path fill="#ffffff" d="m72.98425 27.519686l81.98425 0l0 83.338585l-81.98425 0z" fill-rule="evenodd"/><path stroke="#000000" stroke-width="1.0" stroke-linejoin="round" stroke-linecap="butt" d="m72.98425 27.519686l81.98425 0l0 83.338585l-81.98425 0z" fill-rule="evenodd"/><path fill="#000000" d="m86.828 76.10897l0 -11.78125l-4.40625 0l0 -1.578125l10.578125 0l0 1.578125l-4.40625 0l0 11.78125l-1.765625 0zm7.0782776 0l0 -9.671875l1.46875 0l0 1.46875q0.5625 -1.03125 1.03125 -1.359375q0.484375 -0.328125 1.0625 -0.328125q0.828125 0 1.6875 0.53125l-0.5625 1.515625q-0.609375 -0.359375 -1.203125 -0.359375q-0.546875 0 -0.96875 0.328125q-0.421875 0.328125 -0.609375 0.890625q-0.28125 0.875 -0.28125 1.921875l0 5.0625l-1.625 0zm12.540802 -1.1875q-0.921875 0.765625 -1.765625 1.09375q-0.828125 0.3125 -1.796875 0.3125q-1.59375 0 -2.453125 -0.78125q-0.859375 -0.78125 -0.859375 -1.984375q0 -0.71875 0.328125 -1.296875q0.328125 -0.59375 0.84375 -0.9375q0.53125 -0.359375 1.1875 -0.546875q0.46875 -0.125 1.453125 -0.25q1.984375 -0.234375 2.921875 -0.5625q0.015625 -0.34375 0.015625 -0.421875q0 -1.0 -0.46875 -1.421875q-0.625 -0.546875 -1.875 -0.546875q-1.15625 0 -1.703125 0.40625q-0.546875 0.40625 -0.8125 1.421875l-1.609375 -0.21875q0.21875 -1.015625 0.71875 -1.640625q0.5 -0.640625 1.453125 -0.984375q0.953125 -0.34375 2.1875 -0.34375q1.25 0 2.015625 0.296875q0.78125 0.28125 1.140625 0.734375q0.375 0.4375 0.515625 1.109375q0.078125 0.421875 0.078125 1.515625l0 2.1875q0 2.28125 0.109375 2.890625q0.109375 0.59375 0.40625 1.15625l-1.703125 0q-0.265625 -0.515625 -0.328125 -1.1875zm-0.140625 -3.671875q-0.890625 0.375 -2.671875 0.625q-1.015625 0.140625 -1.4375 0.328125q-0.421875 0.1875 -0.65625 0.53125q-0.21875 0.34375 -0.21875 0.78125q0 0.65625 0.5 1.09375q0.5 0.4375 1.453125 0.4375q0.9375 0 1.671875 -0.40625q0.75 -0.421875 1.09375 -1.140625q0.265625 -0.5625 0.265625 -1.640625l0 -0.609375zm4.203842 -6.609375l0 -1.890625l1.640625 0l0 1.890625l-1.640625 0zm0 11.46875l0 -9.671875l1.640625 0l0 9.671875l-1.640625 0zm4.144821 0l0 -9.671875l1.46875 0l0 1.375q1.0625 -1.59375 3.078125 -1.59375q0.875 0 1.609375 0.3125q0.734375 0.3125 1.09375 0.828125q0.375 0.5 0.515625 1.203125q0.09375 0.453125 0.09375 1.59375l0 5.953125l-1.640625 0l0 -5.890625q0 -1.0 -0.203125 -1.484375q-0.1875 -0.5 -0.671875 -0.796875q-0.484375 -0.296875 -1.140625 -0.296875q-1.046875 0 -1.8125 0.671875q-0.75 0.65625 -0.75 2.515625l0 5.28125l-1.640625 0zm17.000717 -3.109375l1.6875 0.203125q-0.40625 1.484375 -1.484375 2.3125q-1.078125 0.8125 -2.765625 0.8125q-2.125 0 -3.375 -1.296875q-1.234375 -1.3125 -1.234375 -3.671875q0 -2.453125 1.25 -3.796875q1.265625 -1.34375 3.265625 -1.34375q1.9375 0 3.15625 1.328125q1.234375 1.3125 1.234375 3.703125q0 0.15625 0 0.4375l-7.21875 0q0.09375 1.59375 0.90625 2.453125q0.8125 0.84375 2.015625 0.84375q0.90625 0 1.546875 -0.46875q0.640625 -0.484375 1.015625 -1.515625zm-5.390625 -2.65625l5.40625 0q-0.109375 -1.21875 -0.625 -1.828125q-0.78125 -0.953125 -2.03125 -0.953125q-1.125 0 -1.90625 0.765625q-0.765625 0.75 -0.84375 2.015625zm9.125717 5.765625l0 -9.671875l1.46875 0l0 1.46875q0.5625 -1.03125 1.03125 -1.359375q0.484375 -0.328125 1.0625 -0.328125q0.828125 0 1.6875 0.53125l-0.5625 1.515625q-0.609375 -0.359375 -1.203125 -0.359375q-0.546875 0 -0.96875 0.328125q-0.421875 0.328125 -0.609375 0.890625q-0.28125 0.875 -0.28125 1.921875l0 5.0625l-1.625 0z" fill-rule="nonzero"/><path fill="#ffffff" d="m302.44882 32.372704l223.2756 0l0 46.708664l-223.2756 0z" fill-rule="evenodd"/><path stroke="#000000" stroke-width="1.0" stroke-linejoin="round" stroke-linecap="butt" d="m302.44882 32.372704l223.2756 0l0 46.708664l-223.2756 0z" fill-rule="evenodd"/><path fill="#000000" d="m349.37497 62.647034l0 -13.359375l5.921875 0q1.78125 0 2.703125 0.359375q0.9375 0.359375 1.484375 1.28125q0.5625 0.90625 0.5625 2.015625q0 1.40625 -0.921875 2.390625q-0.921875 0.96875 -2.84375 1.234375q0.703125 0.34375 1.078125 0.671875q0.765625 0.703125 1.453125 1.765625l2.328125 3.640625l-2.21875 0l-1.765625 -2.78125q-0.78125 -1.203125 -1.28125 -1.828125q-0.5 -0.640625 -0.90625 -0.890625q-0.390625 -0.265625 -0.796875 -0.359375q-0.296875 -0.078125 -0.984375 -0.078125l-2.046875 0l0 5.9375l-1.765625 0zm1.765625 -7.453125l3.796875 0q1.21875 0 1.890625 -0.25q0.6875 -0.265625 1.046875 -0.8125q0.359375 -0.546875 0.359375 -1.1875q0 -0.953125 -0.6875 -1.5625q-0.6875 -0.609375 -2.1875 -0.609375l-4.21875 0l0 4.421875zm10.863556 2.609375q0 -2.6875 1.484375 -3.96875q1.25 -1.078125 3.046875 -1.078125q2.0 0 3.265625 1.3125q1.265625 1.296875 1.265625 3.609375q0 1.859375 -0.5625 2.9375q-0.5625 1.0Line truncated
|
||||
|
After Width: | Height: | Size: 57 KiB |
@@ -104,7 +104,9 @@ Asynchronous Proximal Policy Optimization (APPO)
|
||||
`[implementation] <https://github.com/ray-project/ray/blob/master/rllib/agents/ppo/appo.py>`__
|
||||
We include an asynchronous variant of Proximal Policy Optimization (PPO) based on the IMPALA architecture. This is similar to IMPALA but using a surrogate policy loss with clipping. Compared to synchronous PPO, APPO is more efficient in wall-clock time due to its use of asynchronous sampling. Using a clipped loss also allows for multiple SGD passes, and therefore the potential for better sample efficiency compared to IMPALA. V-trace can also be enabled to correct for off-policy samples.
|
||||
|
||||
APPO is not always more efficient; it is often better to simply use `PPO <rllib-algorithms.html#proximal-policy-optimization-ppo>`__ or `IMPALA <rllib-algorithms.html#importance-weighted-actor-learner-architecture-impala>`__.
|
||||
.. tip::
|
||||
|
||||
APPO is not always more efficient; it is often better to use `standard PPO <rllib-algorithms.html#proximal-policy-optimization-ppo>`__ or `IMPALA <rllib-algorithms.html#importance-weighted-actor-learner-architecture-impala>`__.
|
||||
|
||||
.. figure:: impala-arch.svg
|
||||
|
||||
@@ -119,6 +121,30 @@ Tuned examples: `PongNoFrameskip-v4 <https://github.com/ray-project/ray/blob/mas
|
||||
:start-after: __sphinx_doc_begin__
|
||||
:end-before: __sphinx_doc_end__
|
||||
|
||||
Decentralized Distributed Proximal Policy Optimization (DD-PPO)
|
||||
---------------------------------------------------------------
|
||||
|pytorch|
|
||||
`[paper] <https://arxiv.org/abs/1911.00357>`__
|
||||
`[implementation] <https://github.com/ray-project/ray/blob/master/rllib/agents/ppo/ddppo.py>`__
|
||||
Unlike APPO or PPO, with DD-PPO policy improvement is no longer done centralized in the trainer process. Instead, gradients are computed remotely on each rollout worker and all-reduced at each mini-batch using `torch distributed <https://pytorch.org/docs/stable/distributed.html>`__. This allows each worker's GPU to be used both for sampling and for training.
|
||||
|
||||
.. tip::
|
||||
|
||||
DD-PPO is best for envs that require GPUs to function, or if you need to scale out SGD to multiple nodes. If you don't meet these requirements, `standard PPO <#proximal-policy-optimization-ppo>`__ will be more efficient.
|
||||
|
||||
.. figure:: ddppo-arch.svg
|
||||
|
||||
DD-PPO architecture (both sampling and learning are done on worker GPUs)
|
||||
|
||||
Tuned examples: `CartPole-v0 <https://github.com/ray-project/ray/blob/master/rllib/tuned_examples/regression_tests/cartpole-ddppo.yaml>`__, `BreakoutNoFrameskip-v4 <https://github.com/ray-project/ray/blob/master/rllib/tuned_examples/atari-ddppo.yaml>`__
|
||||
|
||||
**DDPPO-specific configs** (see also `common configs <rllib-training.html#common-parameters>`__):
|
||||
|
||||
.. literalinclude:: ../../rllib/agents/ppo/ddppo.py
|
||||
:language: python
|
||||
:start-after: __sphinx_doc_begin__
|
||||
:end-before: __sphinx_doc_end__
|
||||
|
||||
Gradient-based
|
||||
~~~~~~~~~~~~~~
|
||||
|
||||
@@ -241,7 +267,11 @@ Proximal Policy Optimization (PPO)
|
||||
----------------------------------
|
||||
|pytorch| |tensorflow|
|
||||
`[paper] <https://arxiv.org/abs/1707.06347>`__ `[implementation] <https://github.com/ray-project/ray/blob/master/rllib/agents/ppo/ppo.py>`__
|
||||
PPO's clipped objective supports multiple SGD passes over the same batch of experiences. RLlib's multi-GPU optimizer pins that data in GPU memory to avoid unnecessary transfers from host memory, substantially improving performance over a naive implementation. RLlib's PPO scales out using multiple workers for experience collection, and also with multiple GPUs for SGD.
|
||||
PPO's clipped objective supports multiple SGD passes over the same batch of experiences. RLlib's multi-GPU optimizer pins that data in GPU memory to avoid unnecessary transfers from host memory, substantially improving performance over a naive implementation. PPO scales out using multiple workers for experience collection, and also to multiple GPUs for SGD.
|
||||
|
||||
.. tip::
|
||||
|
||||
If you need to scale out with GPUs on multiple nodes, consider using `decentralized PPO <#decentralized-distributed-proximal-policy-optimization-dd-ppo>`__.
|
||||
|
||||
.. figure:: ppo-arch.svg
|
||||
|
||||
|
||||
@@ -86,6 +86,8 @@ Algorithms
|
||||
|
||||
- |tensorflow| `Asynchronous Proximal Policy Optimization (APPO) <rllib-algorithms.html#asynchronous-proximal-policy-optimization-appo>`__
|
||||
|
||||
- |pytorch| `Decentralized Distributed Proximal Policy Optimization (DD-PPO) <rllib-algorithms.html#decentralized-distributed-proximal-policy-optimization-dd-ppo>`__
|
||||
|
||||
- |pytorch| `Single-Player AlphaZero (contrib/AlphaZero) <rllib-algorithms.html#single-player-alpha-zero-contrib-alphazero>`__
|
||||
|
||||
* Gradient-based
|
||||
|
||||
@@ -10,7 +10,7 @@ To get started, take a look over the `custom env example <https://github.com/ray
|
||||
RLlib in 60 seconds
|
||||
-------------------
|
||||
|
||||
The following is a whirlwind overview of RLlib. For a more in-depth guide, see also the `full table of contents <rllib-toc.html>`__ and `RLlib blog posts <rllib-examples.html#blog-posts>`__. You may also want to skim the `list of built-in algorithms <rllib-toc.html#algorithms>`__. Look out for the |tensorflow| and |pytorch| icons to see which algorithms are available for each framework.
|
||||
The following is a whirlwind overview of RLlib. For a more in-depth guide, see also the `full table of contents <rllib-toc.html>`__ and `RLlib blog posts <rllib-examples.html#blog-posts>`__. You may also want to skim the `list of built-in algorithms <rllib-toc.html#algorithms>`__. Look out for the |tensorflow| and |pytorch| icons to see which algorithms are `available <rllib-toc.html#algorithms>`__ for each framework.
|
||||
|
||||
Running RLlib
|
||||
~~~~~~~~~~~~~
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from ray.rllib.agents.ppo.ppo import PPOTrainer, DEFAULT_CONFIG
|
||||
from ray.rllib.agents.ppo.appo import APPOTrainer
|
||||
from ray.rllib.agents.ppo.ddppo import DDPPOTrainer
|
||||
|
||||
__all__ = ["APPOTrainer", "PPOTrainer", "DEFAULT_CONFIG"]
|
||||
__all__ = ["APPOTrainer", "DDPPOTrainer", "PPOTrainer", "DEFAULT_CONFIG"]
|
||||
@@ -0,0 +1,87 @@
|
||||
from ray.rllib.agents.ppo import ppo
|
||||
from ray.rllib.agents.trainer import with_base_config
|
||||
from ray.rllib.optimizers import TorchDistributedDataParallelOptimizer
|
||||
"""Decentralized Distributed PPO implementation.
|
||||
|
||||
Unlike APPO or PPO, learning is no longer done centralized in the trainer
|
||||
process. Instead, gradients are computed remotely on each rollout worker and
|
||||
all-reduced to sync them at each mini-batch. This allows each worker's GPU
|
||||
to be used both for sampling and for training.
|
||||
|
||||
DD-PPO should be used if you have envs that require GPUs to function, or have
|
||||
a very large model that cannot be effectively optimized with the GPUs available
|
||||
on a single machine (DD-PPO allows scaling to arbitrary numbers of GPUs across
|
||||
multiple nodes, unlike PPO/APPO which is limited to GPUs on a single node).
|
||||
|
||||
Paper reference: https://arxiv.org/abs/1911.00357
|
||||
Note that unlike the paper, we currently do not implement straggler mitigation.
|
||||
"""
|
||||
|
||||
# yapf: disable
|
||||
# __sphinx_doc_begin__
|
||||
DEFAULT_CONFIG = with_base_config(ppo.DEFAULT_CONFIG, {
|
||||
# During the sampling phase, each rollout worker will collect a batch
|
||||
# `sample_batch_size * num_envs_per_worker` steps in size.
|
||||
"sample_batch_size": 100,
|
||||
# Vectorize the env (should enable by default since each worker has a GPU).
|
||||
"num_envs_per_worker": 5,
|
||||
# During the SGD phase, workers iterate over minibatches of this size.
|
||||
# The effective minibatch size will be `sgd_minibatch_size * num_workers`.
|
||||
"sgd_minibatch_size": 50,
|
||||
# Number of SGD epochs per optimization round.
|
||||
"num_sgd_iter": 10,
|
||||
|
||||
# *** WARNING: configs below are DDPPO overrides over PPO; you
|
||||
# shouldn't need to adjust them. ***
|
||||
"use_pytorch": True, # DDPPO requires PyTorch distributed.
|
||||
"num_gpus": 0, # Learning is no longer done on the driver process, so
|
||||
# giving GPUs to the driver does not make sense!
|
||||
"num_gpus_per_worker": 1, # Each rollout worker gets a GPU.
|
||||
"truncate_episodes": True, # Require evenly sized batches. Otherwise,
|
||||
# collective allreduce could fail.
|
||||
"train_batch_size": -1, # This is auto set based on sample batch size.
|
||||
})
|
||||
# __sphinx_doc_end__
|
||||
# yapf: enable
|
||||
|
||||
|
||||
def validate_config(config):
|
||||
if config["train_batch_size"] == -1:
|
||||
# Auto set.
|
||||
config["train_batch_size"] = (
|
||||
config["sample_batch_size"] * config["num_envs_per_worker"])
|
||||
else:
|
||||
raise ValueError(
|
||||
"Set sample_batch_size instead of train_batch_size for DDPPO.")
|
||||
ppo.validate_config(config)
|
||||
|
||||
|
||||
def make_distributed_allreduce_optimizer(workers, config):
|
||||
if not config["use_pytorch"]:
|
||||
raise ValueError(
|
||||
"Distributed data parallel is only supported for PyTorch")
|
||||
if config["num_gpus"]:
|
||||
raise ValueError(
|
||||
"When using distributed data parallel, you should set "
|
||||
"num_gpus=0 since all optimization "
|
||||
"is happening on workers. Enable GPUs for workers by setting "
|
||||
"num_gpus_per_worker=1.")
|
||||
if config["batch_mode"] != "truncate_episodes":
|
||||
raise ValueError(
|
||||
"Distributed data parallel requires truncate_episodes "
|
||||
"batch mode.")
|
||||
|
||||
return TorchDistributedDataParallelOptimizer(
|
||||
workers,
|
||||
expected_batch_size=config["sample_batch_size"] *
|
||||
config["num_envs_per_worker"],
|
||||
num_sgd_iter=config["num_sgd_iter"],
|
||||
sgd_minibatch_size=config["sgd_minibatch_size"],
|
||||
standardize_fields=["advantages"])
|
||||
|
||||
|
||||
DDPPOTrainer = ppo.PPOTrainer.with_updates(
|
||||
name="DDPPO",
|
||||
default_config=DEFAULT_CONFIG,
|
||||
make_policy_optimizer=make_distributed_allreduce_optimizer,
|
||||
validate_config=validate_config)
|
||||
+1
-31
@@ -3,8 +3,7 @@ import logging
|
||||
from ray.rllib.agents import with_common_config
|
||||
from ray.rllib.agents.ppo.ppo_tf_policy import PPOTFPolicy
|
||||
from ray.rllib.agents.trainer_template import build_trainer
|
||||
from ray.rllib.optimizers import SyncSamplesOptimizer, \
|
||||
LocalMultiGPUOptimizer, TorchDistributedDataParallelOptimizer
|
||||
from ray.rllib.optimizers import SyncSamplesOptimizer, LocalMultiGPUOptimizer
|
||||
from ray.rllib.utils import try_import_tf
|
||||
|
||||
tf = try_import_tf()
|
||||
@@ -69,8 +68,6 @@ DEFAULT_CONFIG = with_common_config({
|
||||
# usually slower, but you might want to try it if you run into issues with
|
||||
# the default optimizer.
|
||||
"simple_optimizer": False,
|
||||
# Use the experimental torch multi-node SGD optimizer.
|
||||
"distributed_data_parallel_optimizer": False,
|
||||
# Use PyTorch as framework?
|
||||
"use_pytorch": False
|
||||
})
|
||||
@@ -79,33 +76,6 @@ DEFAULT_CONFIG = with_common_config({
|
||||
|
||||
|
||||
def choose_policy_optimizer(workers, config):
|
||||
if config["distributed_data_parallel_optimizer"]:
|
||||
if not config["use_pytorch"]:
|
||||
raise ValueError(
|
||||
"Distributed data parallel is only supported for PyTorch")
|
||||
if config["num_gpus"]:
|
||||
raise ValueError(
|
||||
"When using distributed data parallel, you should set "
|
||||
"num_gpus=0 since all optimization "
|
||||
"is happening on workers. Enable GPUs for workers by setting "
|
||||
"num_gpus_per_worker=1.")
|
||||
if config["batch_mode"] != "truncate_episodes":
|
||||
raise ValueError(
|
||||
"Distributed data parallel requires truncate_episodes "
|
||||
"batch mode.")
|
||||
if config["sample_batch_size"] != config["train_batch_size"]:
|
||||
raise ValueError(
|
||||
"Distributed data parallel requires sample_batch_size to be "
|
||||
"equal to train_batch_size. Each worker will sample and learn "
|
||||
"on train_batch_size samples per iteration.")
|
||||
|
||||
return TorchDistributedDataParallelOptimizer(
|
||||
workers,
|
||||
num_sgd_iter=config["num_sgd_iter"],
|
||||
train_batch_size=config["train_batch_size"],
|
||||
sgd_minibatch_size=config["sgd_minibatch_size"],
|
||||
standardize_fields=["advantages"])
|
||||
|
||||
if config["simple_optimizer"]:
|
||||
return SyncSamplesOptimizer(
|
||||
workers,
|
||||
|
||||
@@ -161,9 +161,10 @@ def vf_preds_and_logits_fetches(policy, input_dict, state_batches, model,
|
||||
action_dist):
|
||||
"""Adds value function and logits outputs to experience train_batches."""
|
||||
return {
|
||||
SampleBatch.VF_PREDS: policy.model.value_function(),
|
||||
BEHAVIOUR_LOGITS: policy.model.last_output().numpy(),
|
||||
ACTION_LOGP: action_dist.logp(input_dict[SampleBatch.ACTIONS])
|
||||
SampleBatch.VF_PREDS: policy.model.value_function().cpu().numpy(),
|
||||
BEHAVIOUR_LOGITS: policy.model.last_output().cpu().numpy(),
|
||||
ACTION_LOGP: action_dist.logp(
|
||||
input_dict[SampleBatch.ACTIONS]).cpu().numpy(),
|
||||
}
|
||||
|
||||
|
||||
@@ -187,11 +188,14 @@ class ValueNetworkMixin:
|
||||
|
||||
def value(ob, prev_action, prev_reward, *state):
|
||||
model_out, _ = self.model({
|
||||
SampleBatch.CUR_OBS: torch.Tensor([ob]),
|
||||
SampleBatch.PREV_ACTIONS: torch.Tensor([prev_action]),
|
||||
SampleBatch.PREV_REWARDS: torch.Tensor([prev_reward]),
|
||||
SampleBatch.CUR_OBS: torch.Tensor([ob]).to(self.device),
|
||||
SampleBatch.PREV_ACTIONS: torch.Tensor([prev_action]).to(
|
||||
self.device),
|
||||
SampleBatch.PREV_REWARDS: torch.Tensor([prev_reward]).to(
|
||||
self.device),
|
||||
"is_training": False,
|
||||
}, [torch.Tensor([s]) for s in state], torch.Tensor([1]))
|
||||
}, [torch.Tensor([s]).to(self.device) for s in state],
|
||||
torch.Tensor([1]).to(self.device))
|
||||
return self.model.value_function()[0]
|
||||
|
||||
else:
|
||||
|
||||
@@ -15,6 +15,11 @@ def _import_appo():
|
||||
return ppo.APPOTrainer
|
||||
|
||||
|
||||
def _import_ddppo():
|
||||
from ray.rllib.agents import ppo
|
||||
return ppo.DDPPOTrainer
|
||||
|
||||
|
||||
def _import_qmix():
|
||||
from ray.rllib.agents import qmix
|
||||
return qmix.QMixTrainer
|
||||
@@ -113,6 +118,7 @@ ALGORITHMS = {
|
||||
"QMIX": _import_qmix,
|
||||
"APEX_QMIX": _import_apex_qmix,
|
||||
"APPO": _import_appo,
|
||||
"DDPPO": _import_ddppo,
|
||||
"MARWIL": _import_marwil,
|
||||
}
|
||||
|
||||
|
||||
@@ -625,14 +625,14 @@ class RolloutWorker(EvaluatorInterface):
|
||||
logger.debug("Training out:\n\n{}\n".format(summarize(info_out)))
|
||||
return info_out
|
||||
|
||||
def sample_and_learn(self, train_batch_size, num_sgd_iter,
|
||||
def sample_and_learn(self, expected_batch_size, num_sgd_iter,
|
||||
sgd_minibatch_size, standardize_fields):
|
||||
"""Sample and batch and learn on it.
|
||||
|
||||
This is typically used in combination with distributed allreduce.
|
||||
|
||||
Arguments:
|
||||
train_batch_size (int): Number of samples to learn on.
|
||||
expected_batch_size (int): Expected number of samples to learn on.
|
||||
num_sgd_iter (int): Number of SGD iterations.
|
||||
sgd_minibatch_size (int): SGD minibatch size.
|
||||
standardize_fields (list): List of sample fields to normalize.
|
||||
@@ -642,10 +642,12 @@ class RolloutWorker(EvaluatorInterface):
|
||||
count: number of samples learned on.
|
||||
"""
|
||||
batch = self.sample()
|
||||
assert batch.count == train_batch_size, \
|
||||
(batch.count, "Batch size possibly out of sync between workers")
|
||||
assert batch.count == expected_batch_size, \
|
||||
("Batch size possibly out of sync between workers, expected:",
|
||||
expected_batch_size, "got:", batch.count)
|
||||
logger.info("Executing distributed minibatch SGD "
|
||||
"on batch of size {}".format(batch.count))
|
||||
"with epoch size {}, minibatch size {}".format(
|
||||
batch.count, sgd_minibatch_size))
|
||||
info = do_minibatch_sgd(batch, self.policy_map, self, num_sgd_iter,
|
||||
sgd_minibatch_size, standardize_fields)
|
||||
return info, batch.count
|
||||
|
||||
@@ -13,8 +13,8 @@ class TorchDistributedDataParallelOptimizer(PolicyOptimizer):
|
||||
|
||||
def __init__(self,
|
||||
workers,
|
||||
expected_batch_size,
|
||||
num_sgd_iter=1,
|
||||
train_batch_size=1,
|
||||
sgd_minibatch_size=0,
|
||||
standardize_fields=frozenset([]),
|
||||
keep_local_weights_in_sync=True,
|
||||
@@ -22,11 +22,12 @@ class TorchDistributedDataParallelOptimizer(PolicyOptimizer):
|
||||
PolicyOptimizer.__init__(self, workers)
|
||||
self.learner_stats = {}
|
||||
self.num_sgd_iter = num_sgd_iter
|
||||
self.train_batch_size = train_batch_size
|
||||
self.expected_batch_size = expected_batch_size
|
||||
self.sgd_minibatch_size = sgd_minibatch_size
|
||||
self.standardize_fields = standardize_fields
|
||||
self.keep_local_weights_in_sync = keep_local_weights_in_sync
|
||||
self.update_weights_timer = TimerStat()
|
||||
self.sync_down_timer = TimerStat()
|
||||
self.sync_up_timer = TimerStat()
|
||||
self.learn_timer = TimerStat()
|
||||
|
||||
# Setup the distributed processes.
|
||||
@@ -52,7 +53,7 @@ class TorchDistributedDataParallelOptimizer(PolicyOptimizer):
|
||||
# add too much overhead and handles the case where the user manually
|
||||
# updates the local weights.
|
||||
if self.keep_local_weights_in_sync:
|
||||
with self.update_weights_timer:
|
||||
with self.sync_up_timer:
|
||||
weights = ray.put(self.workers.local_worker().get_weights())
|
||||
for e in self.workers.remote_workers():
|
||||
e.set_weights.remote(weights)
|
||||
@@ -60,7 +61,7 @@ class TorchDistributedDataParallelOptimizer(PolicyOptimizer):
|
||||
with self.learn_timer:
|
||||
results = ray.get([
|
||||
w.sample_and_learn.remote(
|
||||
self.train_batch_size, self.num_sgd_iter,
|
||||
self.expected_batch_size, self.num_sgd_iter,
|
||||
self.sgd_minibatch_size, self.standardize_fields)
|
||||
for w in self.workers.remote_workers()
|
||||
])
|
||||
@@ -87,8 +88,10 @@ class TorchDistributedDataParallelOptimizer(PolicyOptimizer):
|
||||
# Sync down the weights. As with the sync up, this is not really
|
||||
# needed unless the user is reading the local weights.
|
||||
if self.keep_local_weights_in_sync:
|
||||
self.workers.local_worker().set_weights(
|
||||
ray.get(self.workers.remote_workers()[0].get_weights.remote()))
|
||||
with self.sync_down_timer:
|
||||
self.workers.local_worker().set_weights(
|
||||
ray.get(
|
||||
self.workers.remote_workers()[0].get_weights.remote()))
|
||||
|
||||
return self.learner_stats
|
||||
|
||||
@@ -96,8 +99,10 @@ class TorchDistributedDataParallelOptimizer(PolicyOptimizer):
|
||||
def stats(self):
|
||||
return dict(
|
||||
PolicyOptimizer.stats(self), **{
|
||||
"update_weights_time_ms": round(
|
||||
1000 * self.update_weights_timer.mean, 3),
|
||||
"sync_weights_up_time": round(1000 * self.sync_up_timer.mean,
|
||||
3),
|
||||
"sync_weights_down_time": round(
|
||||
1000 * self.sync_down_timer.mean, 3),
|
||||
"learn_time_ms": round(1000 * self.learn_timer.mean, 3),
|
||||
"learner": self.learner_stats,
|
||||
})
|
||||
@@ -108,8 +108,14 @@ class TorchPolicy(Policy):
|
||||
if p.grad is not None:
|
||||
grads.append(p.grad)
|
||||
start = time.time()
|
||||
torch.distributed.all_reduce_coalesced(
|
||||
grads, op=torch.distributed.ReduceOp.SUM)
|
||||
if torch.cuda.is_available():
|
||||
# Sadly, allreduce_coalesced does not work with CUDA yet.
|
||||
for g in grads:
|
||||
torch.distributed.all_reduce(
|
||||
g, op=torch.distributed.ReduceOp.SUM)
|
||||
else:
|
||||
torch.distributed.all_reduce_coalesced(
|
||||
grads, op=torch.distributed.ReduceOp.SUM)
|
||||
for p in self.model.parameters():
|
||||
if p.grad is not None:
|
||||
p.grad /= self.distributed_world_size
|
||||
@@ -208,6 +214,8 @@ class TorchPolicy(Policy):
|
||||
train_batch = UsageTrackingDict(postprocessed_batch)
|
||||
|
||||
def convert(arr):
|
||||
if torch.is_tensor(arr):
|
||||
return arr.to(self.device)
|
||||
tensor = torch.from_numpy(np.asarray(arr))
|
||||
if tensor.dtype == torch.double:
|
||||
tensor = tensor.float()
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
# Basically the same as atari-ppo, but adapted for DDPPO. Note that DDPPO
|
||||
# isn't actually any more efficient on Atari, since the network size is
|
||||
# relatively small and the env doesn't require a GPU.
|
||||
atari-ddppo:
|
||||
env:
|
||||
grid_search:
|
||||
- BreakoutNoFrameskip-v4
|
||||
run: DDPPO
|
||||
config:
|
||||
# Worker config: 10 workers, each of which requires a GPU.
|
||||
num_workers: 10
|
||||
num_gpus_per_worker: 1
|
||||
# Each worker will sample 100 * 5 envs per worker steps = 500 steps
|
||||
# per optimization round. This is 5000 steps summed across workers.
|
||||
sample_batch_size: 100
|
||||
num_envs_per_worker: 5
|
||||
# Each worker will take a minibatch of 50. There are 10 workers total,
|
||||
# so the effective minibatch size will be 500.
|
||||
sgd_minibatch_size: 50
|
||||
num_sgd_iter: 10
|
||||
# Params from standard PPO Atari config:
|
||||
lambda: 0.95
|
||||
kl_coeff: 0.5
|
||||
clip_rewards: True
|
||||
clip_param: 0.1
|
||||
vf_clip_param: 10.0
|
||||
entropy_coeff: 0.01
|
||||
batch_mode: truncate_episodes
|
||||
observation_filter: NoFilter
|
||||
vf_share_layers: true
|
||||
@@ -0,0 +1,8 @@
|
||||
cartpole-ddppo:
|
||||
env: CartPole-v0
|
||||
run: DDPPO
|
||||
stop:
|
||||
episode_reward_mean: 100
|
||||
timesteps_total: 100000
|
||||
config:
|
||||
num_gpus_per_worker: 0
|
||||
@@ -1,14 +0,0 @@
|
||||
cartpole-torch-dist:
|
||||
env: CartPole-v0
|
||||
run: PPO
|
||||
stop:
|
||||
episode_reward_mean: 150
|
||||
timesteps_total: 100000
|
||||
config:
|
||||
num_workers: 2
|
||||
sample_batch_size: 4000
|
||||
train_batch_size: 4000
|
||||
batch_mode: truncate_episodes
|
||||
observation_filter: MeanStdFilter
|
||||
use_pytorch: true
|
||||
distributed_data_parallel_optimizer: true
|
||||
Reference in new issue
Block a user