From 346885068cca0d928e25503f137ea1a827946e96 Mon Sep 17 00:00:00 2001 From: Megan Kawakami Date: Sat, 16 Feb 2019 19:54:14 -0800 Subject: [PATCH] [rllib] add torch pg (#3857) * add torch pg * add torch imports * added torch pg * working torch pg implementation * add pg pytorch * Update a3c.py * Update a3c.py * Update torch_policy_graph.py * Update torch_policy_graph.py --- doc/source/rllib-algorithms.rst | 2 +- doc/source/rllib-models.rst | 2 +- python/ray/rllib/agents/pg/pg.py | 17 +++-- .../rllib/agents/pg/torch_pg_policy_graph.py | 69 +++++++++++++++++++ .../rllib/evaluation/torch_policy_graph.py | 10 ++- test/jenkins_tests/run_multi_node_tests.sh | 7 ++ 6 files changed, 99 insertions(+), 8 deletions(-) create mode 100644 python/ray/rllib/agents/pg/torch_pg_policy_graph.py diff --git a/doc/source/rllib-algorithms.rst b/doc/source/rllib-algorithms.rst index d7042636e..fd07bdc1b 100644 --- a/doc/source/rllib-algorithms.rst +++ b/doc/source/rllib-algorithms.rst @@ -181,7 +181,7 @@ SpaceInvaders 650 1001 1025 Policy Gradients ---------------- -`[paper] `__ `[implementation] `__ We include a vanilla policy gradients implementation as an example algorithm. This is usually outperformed by PPO. +`[paper] `__ `[implementation] `__ We include a vanilla policy gradients implementation as an example algorithm in both TensorFlow and PyTorch. This is usually outperformed by PPO. Tuned examples: `CartPole-v0 `__ diff --git a/doc/source/rllib-models.rst b/doc/source/rllib-models.rst index 5e3628aa3..6f827809f 100644 --- a/doc/source/rllib-models.rst +++ b/doc/source/rllib-models.rst @@ -155,7 +155,7 @@ You can use ``tf.layers.batch_normalization(x, training=input_dict["is_training" Custom Models (PyTorch) ----------------------- -Similarly, you can create and register custom PyTorch models for use with PyTorch-based algorithms (e.g., A2C, QMIX). See these examples of `fully connected `__, `convolutional `__, and `recurrent `__ torch models. +Similarly, you can create and register custom PyTorch models for use with PyTorch-based algorithms (e.g., A2C, PG, QMIX). See these examples of `fully connected `__, `convolutional `__, and `recurrent `__ torch models. .. code-block:: python diff --git a/python/ray/rllib/agents/pg/pg.py b/python/ray/rllib/agents/pg/pg.py index ba525887e..59428f76a 100644 --- a/python/ray/rllib/agents/pg/pg.py +++ b/python/ray/rllib/agents/pg/pg.py @@ -4,6 +4,7 @@ from __future__ import print_function from ray.rllib.agents.agent import Agent, with_common_config from ray.rllib.agents.pg.pg_policy_graph import PGPolicyGraph + from ray.rllib.optimizers import SyncSamplesOptimizer from ray.rllib.utils.annotations import override @@ -14,6 +15,8 @@ DEFAULT_CONFIG = with_common_config({ "num_workers": 0, # Learning rate "lr": 0.0004, + # Use PyTorch as backend + "use_pytorch": False, }) # __sphinx_doc_end__ # yapf: enable @@ -32,10 +35,16 @@ class PGAgent(Agent): @override(Agent) def _init(self): + if self.config["use_pytorch"]: + from ray.rllib.agents.pg.torch_pg_policy_graph import \ + PGTorchPolicyGraph + policy_cls = PGTorchPolicyGraph + else: + policy_cls = self._policy_graph self.local_evaluator = self.make_local_evaluator( - self.env_creator, self._policy_graph) + self.env_creator, policy_cls) self.remote_evaluators = self.make_remote_evaluators( - self.env_creator, self._policy_graph, self.config["num_workers"]) + self.env_creator, policy_cls, self.config["num_workers"]) optimizer_config = dict( self.config["optimizer"], **{"train_batch_size": self.config["train_batch_size"]}) @@ -48,6 +57,6 @@ class PGAgent(Agent): self.optimizer.step() result = self.optimizer.collect_metrics( self.config["collect_metrics_timeout"]) - result.update(timesteps_this_iter=self.optimizer.num_steps_sampled - - prev_steps) + result.update( + timesteps_this_iter=self.optimizer.num_steps_sampled - prev_steps) return result diff --git a/python/ray/rllib/agents/pg/torch_pg_policy_graph.py b/python/ray/rllib/agents/pg/torch_pg_policy_graph.py new file mode 100644 index 000000000..a818013d4 --- /dev/null +++ b/python/ray/rllib/agents/pg/torch_pg_policy_graph.py @@ -0,0 +1,69 @@ +from __future__ import absolute_import +from __future__ import division +from __future__ import print_function + +import torch +import torch.nn.functional as F +from torch import nn + +import ray +from ray.rllib.models.catalog import ModelCatalog +from ray.rllib.evaluation.postprocessing import compute_advantages +from ray.rllib.evaluation.policy_graph import PolicyGraph +from ray.rllib.evaluation.torch_policy_graph import TorchPolicyGraph +from ray.rllib.utils.annotations import override + + +class PGLoss(nn.Module): + def __init__(self, policy_model): + nn.Module.__init__(self) + self.policy_model = policy_model + + def forward(self, observations, actions, advantages): + logits, _, values, _ = self.policy_model({"obs": observations}, []) + log_probs = F.log_softmax(logits, dim=1) + probs = F.softmax(logits, dim=1) + action_log_probs = log_probs.gather(1, actions.view(-1, 1)) + pi_err = -advantages.dot(action_log_probs.reshape(-1)) + return pi_err + + +class PGTorchPolicyGraph(TorchPolicyGraph): + def __init__(self, obs_space, action_space, config): + config = dict(ray.rllib.agents.a3c.a3c.DEFAULT_CONFIG, **config) + self.config = config + _, self.logit_dim = ModelCatalog.get_action_dist( + action_space, self.config["model"]) + self.model = ModelCatalog.get_torch_model(obs_space, self.logit_dim, + self.config["model"]) + loss = PGLoss(self.model) + + TorchPolicyGraph.__init__( + self, + obs_space, + action_space, + self.model, + loss, + loss_inputs=["obs", "actions", "advantages"]) + + @override(TorchPolicyGraph) + def extra_action_out(self, model_out): + return {"vf_preds": model_out[2].numpy()} + + @override(TorchPolicyGraph) + def optimizer(self): + return torch.optim.Adam(self.model.parameters(), lr=self.config["lr"]) + + @override(PolicyGraph) + def postprocess_trajectory(self, + sample_batch, + other_agent_batches=None, + episode=None): + return compute_advantages( + sample_batch, 0.0, self.config["gamma"], use_gae=False) + + def _value(self, obs): + with self.lock: + obs = torch.from_numpy(obs).float().unsqueeze(0) + _, _, vf, _ = self.model({"obs": obs}, []) + return vf.detach().numpy().squeeze() diff --git a/python/ray/rllib/evaluation/torch_policy_graph.py b/python/ray/rllib/evaluation/torch_policy_graph.py index b7e0f4a07..be55dde86 100644 --- a/python/ray/rllib/evaluation/torch_policy_graph.py +++ b/python/ray/rllib/evaluation/torch_policy_graph.py @@ -85,14 +85,20 @@ class TorchPolicyGraph(PolicyGraph): loss_out.backward() # Note that return values are just references; # calling zero_grad will modify the values - grads = [p.grad.data.numpy() for p in self._model.parameters()] + grads = [] + for p in self._model.parameters(): + if p.grad is not None: + grads.append(p.grad.data.numpy()) + else: + grads.append(None) return grads, {} @override(PolicyGraph) def apply_gradients(self, gradients): with self.lock: for g, p in zip(gradients, self._model.parameters()): - p.grad = torch.from_numpy(g) + if g is not None: + p.grad = torch.from_numpy(g) self._optimizer.step() return {} diff --git a/test/jenkins_tests/run_multi_node_tests.sh b/test/jenkins_tests/run_multi_node_tests.sh index b513a3148..c75d76e1d 100755 --- a/test/jenkins_tests/run_multi_node_tests.sh +++ b/test/jenkins_tests/run_multi_node_tests.sh @@ -167,6 +167,13 @@ docker run --rm --shm-size=${SHM_SIZE} --memory=${MEMORY_SIZE} $DOCKER_SHA \ --stop '{"training_iteration": 2}' \ --config '{"sample_batch_size": 500, "num_workers": 1}' +docker run --rm --shm-size=${SHM_SIZE} --memory=${MEMORY_SIZE} $DOCKER_SHA \ + python /ray/python/ray/rllib/train.py \ + --env CartPole-v0 \ + --run PG \ + --stop '{"training_iteration": 2}' \ + --config '{"sample_batch_size": 500, "use_pytorch": true}' + docker run --rm --shm-size=${SHM_SIZE} --memory=${MEMORY_SIZE} $DOCKER_SHA \ python /ray/python/ray/rllib/train.py \ --env CartPole-v0 \