From 499ad5fbe4d0326f8970f9dd10a88d8e96278f74 Mon Sep 17 00:00:00 2001 From: Sven Mika Date: Thu, 23 Apr 2020 09:11:12 +0200 Subject: [PATCH] [RLlib] PyTorch version of APPO. (#8120) - Translate all vtrace functionality to torch and added torch to the framework_iterator-loop in all existing vtrace test cases. - Add learning test cases for APPO torch (both w/ and w/o v-trace). - Add quick compilation tests for APPO (tf and torch, v-trace and no v-trace). --- doc/source/rllib-algorithms.rst | 2 +- doc/source/rllib-toc.rst | 2 +- rllib/BUILD | 16 + rllib/agents/impala/impala.py | 4 +- rllib/agents/impala/tests/test_vtrace.py | 367 ++++++++++------ .../agents/impala/{vtrace.py => vtrace_tf.py} | 46 +- .../{vtrace_policy.py => vtrace_tf_policy.py} | 2 +- rllib/agents/impala/vtrace_torch.py | 344 +++++++++++++++ rllib/agents/impala/vtrace_torch_policy.py | 57 +++ rllib/agents/ppo/appo.py | 18 +- .../ppo/{appo_policy.py => appo_tf_policy.py} | 10 +- rllib/agents/ppo/appo_torch_policy.py | 411 ++++++++++++++++++ rllib/agents/ppo/tests/test_appo.py | 42 ++ rllib/policy/torch_policy_template.py | 30 +- ...appo-vtrace.yaml => pendulum-appo-tf.yaml} | 11 +- rllib/tuned_examples/pendulum-appo-torch.yaml | 21 + ...rtpole-appo.yaml => cartpole-appo-tf.yaml} | 7 +- .../regression_tests/cartpole-appo-torch.yaml | 14 + ...race.yaml => cartpole-appo-vtrace-tf.yaml} | 7 +- .../cartpole-appo-vtrace-torch.yaml | 14 + rllib/utils/torch_ops.py | 18 + 21 files changed, 1240 insertions(+), 203 deletions(-) rename rllib/agents/impala/{vtrace.py => vtrace_tf.py} (94%) rename rllib/agents/impala/{vtrace_policy.py => vtrace_tf_policy.py} (99%) create mode 100644 rllib/agents/impala/vtrace_torch.py create mode 100644 rllib/agents/impala/vtrace_torch_policy.py rename rllib/agents/ppo/{appo_policy.py => appo_tf_policy.py} (97%) create mode 100644 rllib/agents/ppo/appo_torch_policy.py create mode 100644 rllib/agents/ppo/tests/test_appo.py rename rllib/tuned_examples/{regression_tests/pendulum-appo-vtrace.yaml => pendulum-appo-tf.yaml} (61%) create mode 100644 rllib/tuned_examples/pendulum-appo-torch.yaml rename rllib/tuned_examples/regression_tests/{cartpole-appo.yaml => cartpole-appo-tf.yaml} (66%) create mode 100644 rllib/tuned_examples/regression_tests/cartpole-appo-torch.yaml rename rllib/tuned_examples/regression_tests/{cartpole-appo-vtrace.yaml => cartpole-appo-vtrace-tf.yaml} (64%) create mode 100644 rllib/tuned_examples/regression_tests/cartpole-appo-vtrace-torch.yaml diff --git a/doc/source/rllib-algorithms.rst b/doc/source/rllib-algorithms.rst index be4eef99c..9e490facd 100644 --- a/doc/source/rllib-algorithms.rst +++ b/doc/source/rllib-algorithms.rst @@ -143,7 +143,7 @@ SpaceInvaders 843 ~300 Asynchronous Proximal Policy Optimization (APPO) ------------------------------------------------ -|tensorflow| +|pytorch| |tensorflow| `[paper] `__ `[implementation] `__ 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. diff --git a/doc/source/rllib-toc.rst b/doc/source/rllib-toc.rst index 1df75a0a2..0c1bf6a87 100644 --- a/doc/source/rllib-toc.rst +++ b/doc/source/rllib-toc.rst @@ -91,7 +91,7 @@ Algorithms - |tensorflow| :ref:`Importance Weighted Actor-Learner Architecture (IMPALA) ` - - |tensorflow| :ref:`Asynchronous Proximal Policy Optimization (APPO) ` + - |pytorch| |tensorflow| :ref:`Asynchronous Proximal Policy Optimization (APPO) ` - |pytorch| :ref:`Decentralized Distributed Proximal Policy Optimization (DD-PPO) ` diff --git a/rllib/BUILD b/rllib/BUILD index ded40dc77..6c6704acc 100644 --- a/rllib/BUILD +++ b/rllib/BUILD @@ -140,6 +140,22 @@ py_test( "agents/ppo/tests/test.py"] # TODO(sven): Move down once PR 6889 merged ) +# DDPPO +py_test( + name = "test_ddppo", + tags = ["agents_dir"], + size = "small", + srcs = ["agents/ppo/tests/test_ddppo.py"] +) + +# APPO +py_test( + name = "test_appo", + tags = ["agents_dir"], + size = "medium", + srcs = ["agents/ppo/tests/test_appo.py"] +) + # SAC py_test( name = "test_sac", diff --git a/rllib/agents/impala/impala.py b/rllib/agents/impala/impala.py index e88043c04..193bddd60 100644 --- a/rllib/agents/impala/impala.py +++ b/rllib/agents/impala/impala.py @@ -1,5 +1,5 @@ from ray.rllib.agents.a3c.a3c_tf_policy import A3CTFPolicy -from ray.rllib.agents.impala.vtrace_policy import VTraceTFPolicy +from ray.rllib.agents.impala.vtrace_tf_policy import VTraceTFPolicy from ray.rllib.agents.trainer import Trainer, with_common_config from ray.rllib.agents.trainer_template import build_trainer from ray.rllib.optimizers import AsyncSamplesOptimizer @@ -11,7 +11,7 @@ from ray.tune.resources import Resources # yapf: disable # __sphinx_doc_begin__ DEFAULT_CONFIG = with_common_config({ - # V-trace params (see vtrace.py). + # V-trace params (see vtrace_tf.py). "vtrace": True, "vtrace_clip_rho_threshold": 1.0, "vtrace_clip_pg_rho_threshold": 1.0, diff --git a/rllib/agents/impala/tests/test_vtrace.py b/rllib/agents/impala/tests/test_vtrace.py index 41df02723..3694887a2 100644 --- a/rllib/agents/impala/tests/test_vtrace.py +++ b/rllib/agents/impala/tests/test_vtrace.py @@ -20,22 +20,21 @@ Importance Weighted Actor-Learner Architectures" by Espeholt, Soyer, Munos et al. """ -from absl.testing import parameterized +from gym.spaces import Box import numpy as np +import unittest -from ray.rllib.utils import try_import_tf -import ray.rllib.agents.impala.vtrace as vtrace +from ray.rllib.agents.impala import vtrace_tf as vtrace_tf +from ray.rllib.agents.impala import vtrace_torch as vtrace_torch +from ray.rllib.utils.framework import try_import_tf, try_import_torch from ray.rllib.utils.numpy import softmax +from ray.rllib.utils.test_utils import check, framework_iterator tf = try_import_tf() +torch, nn = try_import_torch() -def _shaped_arange(*shape): - """Runs np.arange, converts to float and reshapes.""" - return np.arange(np.prod(shape), dtype=np.float32).reshape(*shape) - - -def _ground_truth_calculation(discounts, log_rhos, rewards, values, +def _ground_truth_calculation(vtrace, discounts, log_rhos, rewards, values, bootstrap_value, clip_rho_threshold, clip_pg_rho_threshold): """Calculates the ground truth for V-trace in Python/Numpy.""" @@ -63,13 +62,14 @@ def _ground_truth_calculation(discounts, log_rhos, rewards, values, # of the paper is inclusive of the `t-1`, but Python is exclusive. # Also note that np.prod([]) == 1. values_t_plus_1 = np.concatenate( - [values, bootstrap_value[None, :]], axis=0) + [values[1:], bootstrap_value[None, :]], axis=0) for s in range(seq_len): v_s = np.copy(values[s]) # Very important copy. for t in range(s, seq_len): - v_s += (np.prod(discounts[s:t], axis=0) * np.prod(cs[s:t], axis=0) - * clipped_rhos[t] * (rewards[t] + discounts[t] * - values_t_plus_1[t + 1] - values[t])) + v_s += ( + np.prod(discounts[s:t], axis=0) * np.prod( + cs[s:t], axis=0) * clipped_rhos[t] * + (rewards[t] + discounts[t] * values_t_plus_1[t] - values[t])) vs.append(v_s) vs = np.stack(vs, axis=0) pg_advantages = (clipped_pg_rhos * (rewards + discounts * np.concatenate( @@ -78,185 +78,262 @@ def _ground_truth_calculation(discounts, log_rhos, rewards, values, return vtrace.VTraceReturns(vs=vs, pg_advantages=pg_advantages) -class LogProbsFromLogitsAndActionsTest(tf.test.TestCase, - parameterized.TestCase): - @parameterized.named_parameters(("Batch1", 1), ("Batch2", 2)) - def test_log_probs_from_logits_and_actions(self, batch_size): +class LogProbsFromLogitsAndActionsTest(unittest.TestCase): + def test_log_probs_from_logits_and_actions(self): """Tests log_probs_from_logits_and_actions.""" seq_len = 7 num_actions = 3 + batch_size = 4 - policy_logits = _shaped_arange(seq_len, batch_size, num_actions) + 10 - actions = np.random.randint( - 0, num_actions - 1, size=(seq_len, batch_size), dtype=np.int32) + for fw, sess in framework_iterator( + frameworks=("torch", "tf"), session=True): + vtrace = vtrace_tf if fw == "tf" else vtrace_torch + policy_logits = Box(-1.0, 1.0, (seq_len, batch_size, num_actions), + np.float32).sample() + actions = np.random.randint( + 0, num_actions - 1, size=(seq_len, batch_size), dtype=np.int32) - action_log_probs_tensor = vtrace.log_probs_from_logits_and_actions( - policy_logits, actions) + if fw == "torch": + action_log_probs_tensor = \ + vtrace.log_probs_from_logits_and_actions( + torch.from_numpy(policy_logits), + torch.from_numpy(actions)) + else: + action_log_probs_tensor = \ + vtrace.log_probs_from_logits_and_actions( + policy_logits, actions) - # Ground Truth - # Using broadcasting to create a mask that indexes action logits - action_index_mask = actions[..., None] == np.arange(num_actions) + # Ground Truth + # Using broadcasting to create a mask that indexes action logits + action_index_mask = actions[..., None] == np.arange(num_actions) - def index_with_mask(array, mask): - return array[mask].reshape(*array.shape[:-1]) + def index_with_mask(array, mask): + return array[mask].reshape(*array.shape[:-1]) - # Note: Normally log(softmax) is not a good idea because it's not - # numerically stable. However, in this test we have well-behaved - # values. - ground_truth_v = index_with_mask( - np.log(softmax(policy_logits)), action_index_mask) + # Note: Normally log(softmax) is not a good idea because it's not + # numerically stable. However, in this test we have well-behaved + # values. + ground_truth_v = index_with_mask( + np.log(softmax(policy_logits)), action_index_mask) - with self.test_session() as session: - self.assertAllClose(ground_truth_v, - session.run(action_log_probs_tensor)) + if sess: + action_log_probs_tensor = sess.run(action_log_probs_tensor) + check(action_log_probs_tensor, ground_truth_v) -class VtraceTest(tf.test.TestCase, parameterized.TestCase): - @parameterized.named_parameters(("Batch1", 1), ("Batch5", 5)) - def test_vtrace(self, batch_size): +class VtraceTest(unittest.TestCase): + def test_vtrace(self): """Tests V-trace against ground truth data calculated in python.""" seq_len = 5 + batch_size = 10 # Create log_rhos such that rho will span from near-zero to above the # clipping thresholds. In particular, calculate log_rhos in # [-2.5, 2.5), # so that rho is in approx [0.08, 12.2). - log_rhos = _shaped_arange(seq_len, batch_size) / (batch_size * seq_len) + space_w_time = Box(-1.0, 1.0, (seq_len, batch_size), np.float32) + space_only_batch = Box(-1.0, 1.0, (batch_size, ), np.float32) + log_rhos = space_w_time.sample() / (batch_size * seq_len) log_rhos = 5 * (log_rhos - 0.5) # [0.0, 1.0) -> [-2.5, 2.5). values = { "log_rhos": log_rhos, # T, B where B_i: [0.9 / (i+1)] * T "discounts": np.array([[0.9 / (b + 1) for b in range(batch_size)] for _ in range(seq_len)]), - "rewards": _shaped_arange(seq_len, batch_size), - "values": _shaped_arange(seq_len, batch_size) / batch_size, - "bootstrap_value": _shaped_arange(batch_size) + 1.0, + "rewards": space_w_time.sample(), + "values": space_w_time.sample() / batch_size, + "bootstrap_value": space_only_batch.sample() + 1.0, "clip_rho_threshold": 3.7, "clip_pg_rho_threshold": 2.2, } - output = vtrace.from_importance_weights(**values) + for fw, sess in framework_iterator( + frameworks=("torch", "tf"), session=True): + vtrace = vtrace_tf if fw == "tf" else vtrace_torch + output = vtrace.from_importance_weights(**values) + if sess: + output = sess.run(output) - with self.test_session() as session: - output_v = session.run(output) + ground_truth_v = _ground_truth_calculation(vtrace, **values) + check(output, ground_truth_v) - ground_truth_v = _ground_truth_calculation(**values) - for a, b in zip(ground_truth_v, output_v): - self.assertAllClose(a, b) - - @parameterized.named_parameters(("Batch1", 1), ("Batch2", 2)) - def test_vtrace_from_logits(self, batch_size): + def test_vtrace_from_logits(self): """Tests V-trace calculated from logits.""" seq_len = 5 + batch_size = 15 num_actions = 3 clip_rho_threshold = None # No clipping. clip_pg_rho_threshold = None # No clipping. + space = Box(-1.0, 1.0, (seq_len, batch_size, num_actions)) + action_space = Box( + 0, num_actions - 1, ( + seq_len, + batch_size, + ), dtype=np.int32) + space_w_time = Box(-1.0, 1.0, ( + seq_len, + batch_size, + )) + space_only_batch = Box(-1.0, 1.0, (batch_size, )) - # Intentionally leaving shapes unspecified to test if V-trace can - # deal with that. - placeholders = { - # T, B, NUM_ACTIONS - "behaviour_policy_logits": tf.placeholder( - dtype=tf.float32, shape=[None, None, None]), - # T, B, NUM_ACTIONS - "target_policy_logits": tf.placeholder( - dtype=tf.float32, shape=[None, None, None]), - "actions": tf.placeholder(dtype=tf.int32, shape=[None, None]), - "discounts": tf.placeholder(dtype=tf.float32, shape=[None, None]), - "rewards": tf.placeholder(dtype=tf.float32, shape=[None, None]), - "values": tf.placeholder(dtype=tf.float32, shape=[None, None]), - "bootstrap_value": tf.placeholder(dtype=tf.float32, shape=[None]), - } + for fw, sess in framework_iterator( + frameworks=("torch", "tf"), session=True): + vtrace = vtrace_tf if fw == "tf" else vtrace_torch - from_logits_output = vtrace.from_logits( - clip_rho_threshold=clip_rho_threshold, - clip_pg_rho_threshold=clip_pg_rho_threshold, - **placeholders) + if fw == "tf": + # Intentionally leaving shapes unspecified to test if V-trace + # can deal with that. + inputs_ = { + # T, B, NUM_ACTIONS + "behaviour_policy_logits": tf.placeholder( + dtype=tf.float32, shape=[None, None, None]), + # T, B, NUM_ACTIONS + "target_policy_logits": tf.placeholder( + dtype=tf.float32, shape=[None, None, None]), + "actions": tf.placeholder( + dtype=tf.int32, shape=[None, None]), + "discounts": tf.placeholder( + dtype=tf.float32, shape=[None, None]), + "rewards": tf.placeholder( + dtype=tf.float32, shape=[None, None]), + "values": tf.placeholder( + dtype=tf.float32, shape=[None, None]), + "bootstrap_value": tf.placeholder( + dtype=tf.float32, shape=[None]), + } + else: + inputs_ = { + # T, B, NUM_ACTIONS + "behaviour_policy_logits": space.sample(), + # T, B, NUM_ACTIONS + "target_policy_logits": space.sample(), + "actions": action_space.sample(), + "discounts": space_w_time.sample(), + "rewards": space_w_time.sample(), + "values": space_w_time.sample(), + "bootstrap_value": space_only_batch.sample(), + } + from_logits_output = vtrace.from_logits( + clip_rho_threshold=clip_rho_threshold, + clip_pg_rho_threshold=clip_pg_rho_threshold, + **inputs_) - target_log_probs = vtrace.log_probs_from_logits_and_actions( - placeholders["target_policy_logits"], placeholders["actions"]) - behaviour_log_probs = vtrace.log_probs_from_logits_and_actions( - placeholders["behaviour_policy_logits"], placeholders["actions"]) - log_rhos = target_log_probs - behaviour_log_probs - ground_truth = (log_rhos, behaviour_log_probs, target_log_probs) + if fw == "tf": + target_log_probs = vtrace.log_probs_from_logits_and_actions( + inputs_["target_policy_logits"], inputs_["actions"]) + behaviour_log_probs = vtrace.log_probs_from_logits_and_actions( + inputs_["behaviour_policy_logits"], inputs_["actions"]) + else: + target_log_probs = vtrace.log_probs_from_logits_and_actions( + torch.from_numpy(inputs_["target_policy_logits"]), + torch.from_numpy(inputs_["actions"])) + behaviour_log_probs = vtrace.log_probs_from_logits_and_actions( + torch.from_numpy(inputs_["behaviour_policy_logits"]), + torch.from_numpy(inputs_["actions"])) + log_rhos = target_log_probs - behaviour_log_probs + ground_truth = (log_rhos, behaviour_log_probs, target_log_probs) - values = { - "behaviour_policy_logits": _shaped_arange(seq_len, batch_size, - num_actions), - "target_policy_logits": _shaped_arange(seq_len, batch_size, - num_actions), - "actions": np.random.randint( - 0, num_actions - 1, size=(seq_len, batch_size)), - "discounts": np.array( # T, B where B_i: [0.9 / (i+1)] * T - [[0.9 / (b + 1) for b in range(batch_size)] - for _ in range(seq_len)]), - "rewards": _shaped_arange(seq_len, batch_size), - "values": _shaped_arange(seq_len, batch_size) / batch_size, - "bootstrap_value": _shaped_arange(batch_size) + 1.0, # B - } + if sess: + values = { + "behaviour_policy_logits": space.sample(), + "target_policy_logits": space.sample(), + "actions": action_space.sample(), + "discounts": space_w_time.sample(), + "rewards": space_w_time.sample(), + "values": space_w_time.sample() / batch_size, + "bootstrap_value": space_only_batch.sample() + 1.0, + } + feed_dict = {inputs_[k]: v for k, v in values.items()} + from_logits_output = sess.run( + from_logits_output, feed_dict=feed_dict) + log_rhos, behaviour_log_probs, target_log_probs = sess.run( + ground_truth, feed_dict=feed_dict) - feed_dict = {placeholders[k]: v for k, v in values.items()} - with self.test_session() as session: - from_logits_output_v = session.run( - from_logits_output, feed_dict=feed_dict) - (ground_truth_log_rhos, ground_truth_behaviour_action_log_probs, - ground_truth_target_action_log_probs) = session.run( - ground_truth, feed_dict=feed_dict) + # Calculate V-trace using the ground truth logits. + from_iw = vtrace.from_importance_weights( + log_rhos=log_rhos, + discounts=values["discounts"], + rewards=values["rewards"], + values=values["values"], + bootstrap_value=values["bootstrap_value"], + clip_rho_threshold=clip_rho_threshold, + clip_pg_rho_threshold=clip_pg_rho_threshold) + from_iw = sess.run(from_iw) + else: + from_iw = vtrace.from_importance_weights( + log_rhos=log_rhos, + discounts=inputs_["discounts"], + rewards=inputs_["rewards"], + values=inputs_["values"], + bootstrap_value=inputs_["bootstrap_value"], + clip_rho_threshold=clip_rho_threshold, + clip_pg_rho_threshold=clip_pg_rho_threshold) - # Calculate V-trace using the ground truth logits. - from_iw = vtrace.from_importance_weights( - log_rhos=ground_truth_log_rhos, - discounts=values["discounts"], - rewards=values["rewards"], - values=values["values"], - bootstrap_value=values["bootstrap_value"], - clip_rho_threshold=clip_rho_threshold, - clip_pg_rho_threshold=clip_pg_rho_threshold) - - with self.test_session() as session: - from_iw_v = session.run(from_iw) - - self.assertAllClose(from_iw_v.vs, from_logits_output_v.vs) - self.assertAllClose(from_iw_v.pg_advantages, - from_logits_output_v.pg_advantages) - self.assertAllClose(ground_truth_behaviour_action_log_probs, - from_logits_output_v.behaviour_action_log_probs) - self.assertAllClose(ground_truth_target_action_log_probs, - from_logits_output_v.target_action_log_probs) - self.assertAllClose(ground_truth_log_rhos, - from_logits_output_v.log_rhos) + check(from_iw.vs, from_logits_output.vs) + check(from_iw.pg_advantages, from_logits_output.pg_advantages) + check(behaviour_log_probs, + from_logits_output.behaviour_action_log_probs) + check(target_log_probs, from_logits_output.target_action_log_probs) + check(log_rhos, from_logits_output.log_rhos) def test_higher_rank_inputs_for_importance_weights(self): """Checks support for additional dimensions in inputs.""" - placeholders = { - "log_rhos": tf.placeholder( - dtype=tf.float32, shape=[None, None, 1]), - "discounts": tf.placeholder( - dtype=tf.float32, shape=[None, None, 1]), - "rewards": tf.placeholder( - dtype=tf.float32, shape=[None, None, 42]), - "values": tf.placeholder(dtype=tf.float32, shape=[None, None, 42]), - "bootstrap_value": tf.placeholder( - dtype=tf.float32, shape=[None, 42]) - } - output = vtrace.from_importance_weights(**placeholders) - self.assertEqual(output.vs.shape.as_list()[-1], 42) + for fw in framework_iterator(frameworks=("torch", "tf"), session=True): + vtrace = vtrace_tf if fw == "tf" else vtrace_torch + if fw == "tf": + inputs_ = { + "log_rhos": tf.placeholder( + dtype=tf.float32, shape=[None, None, 1]), + "discounts": tf.placeholder( + dtype=tf.float32, shape=[None, None, 1]), + "rewards": tf.placeholder( + dtype=tf.float32, shape=[None, None, 42]), + "values": tf.placeholder( + dtype=tf.float32, shape=[None, None, 42]), + "bootstrap_value": tf.placeholder( + dtype=tf.float32, shape=[None, 42]) + } + else: + inputs_ = { + "log_rhos": Box(-1.0, 1.0, (8, 10, 1)).sample(), + "discounts": Box(-1.0, 1.0, (8, 10, 1)).sample(), + "rewards": Box(-1.0, 1.0, (8, 10, 42)).sample(), + "values": Box(-1.0, 1.0, (8, 10, 42)).sample(), + "bootstrap_value": Box(-1.0, 1.0, (10, 42)).sample() + } + output = vtrace.from_importance_weights(**inputs_) + check(int(output.vs.shape[-1]), 42) def test_inconsistent_rank_inputs_for_importance_weights(self): """Test one of many possible errors in shape of inputs.""" - placeholders = { - "log_rhos": tf.placeholder( - dtype=tf.float32, shape=[None, None, 1]), - "discounts": tf.placeholder( - dtype=tf.float32, shape=[None, None, 1]), - "rewards": tf.placeholder( - dtype=tf.float32, shape=[None, None, 42]), - "values": tf.placeholder(dtype=tf.float32, shape=[None, None, 42]), - # Should be [None, 42]. - "bootstrap_value": tf.placeholder(dtype=tf.float32, shape=[None]) - } - with self.assertRaisesRegexp(ValueError, "must have rank 2"): - vtrace.from_importance_weights(**placeholders) + for fw in framework_iterator(frameworks=("torch", "tf"), session=True): + vtrace = vtrace_tf if fw == "tf" else vtrace_torch + if fw == "tf": + inputs_ = { + "log_rhos": tf.placeholder( + dtype=tf.float32, shape=[None, None, 1]), + "discounts": tf.placeholder( + dtype=tf.float32, shape=[None, None, 1]), + "rewards": tf.placeholder( + dtype=tf.float32, shape=[None, None, 42]), + "values": tf.placeholder( + dtype=tf.float32, shape=[None, None, 42]), + # Should be [None, 42]. + "bootstrap_value": tf.placeholder( + dtype=tf.float32, shape=[None]) + } + else: + inputs_ = { + "log_rhos": Box(-1.0, 1.0, (7, 15, 1)).sample(), + "discounts": Box(-1.0, 1.0, (7, 15, 1)).sample(), + "rewards": Box(-1.0, 1.0, (7, 15, 42)).sample(), + "values": Box(-1.0, 1.0, (7, 15, 42)).sample(), + # Should be [15, 42]. + "bootstrap_value": Box(-1.0, 1.0, (7, )).sample() + } + with self.assertRaisesRegexp((ValueError, AssertionError), + "must have rank 2"): + vtrace.from_importance_weights(**inputs_) if __name__ == "__main__": diff --git a/rllib/agents/impala/vtrace.py b/rllib/agents/impala/vtrace_tf.py similarity index 94% rename from rllib/agents/impala/vtrace.py rename to rllib/agents/impala/vtrace_tf.py index 7d57814f9..9d3e113f0 100644 --- a/rllib/agents/impala/vtrace.py +++ b/rllib/agents/impala/vtrace_tf.py @@ -55,36 +55,26 @@ def multi_log_probs_from_logits_and_actions(policy_logits, actions, dist_class, model): """Computes action log-probs from policy logits and actions. - In the notation used throughout documentation and comments, T refers to the - time dimension ranging from 0 to T-1. B refers to the batch size and - ACTION_SPACE refers to the list of numbers each representing a number of - actions. + In the notation used throughout documentation and comments, T refers to the + time dimension ranging from 0 to T-1. B refers to the batch size and + ACTION_SPACE refers to the list of numbers each representing a number of + actions. - Args: - policy_logits: A list with length of ACTION_SPACE of float32 - tensors of shapes - [T, B, ACTION_SPACE[0]], - ..., - [T, B, ACTION_SPACE[-1]] - with un-normalized log-probabilities parameterizing a softmax policy. - actions: A list with length of ACTION_SPACE of - tensors of shapes - [T, B, ...], - ..., - [T, B, ...] - with actions. - dist_class: Python class of the action distribution - - Returns: - A list with length of ACTION_SPACE of float32 - tensors of shapes - [T, B], - ..., - [T, B] - corresponding to the sampling log probability - of the chosen action w.r.t. the policy. - """ + Args: + policy_logits: A list with length of ACTION_SPACE of float32 + tensors of shapes [T, B, ACTION_SPACE[0]], ..., + [T, B, ACTION_SPACE[-1]] with un-normalized log-probabilities + parameterizing a softmax policy. + actions: A list with length of ACTION_SPACE of tensors of shapes + [T, B, ...], ..., [T, B, ...] + with actions. + dist_class: Python class of the action distribution. + Returns: + A list with length of ACTION_SPACE of float32 tensors of shapes + [T, B], ..., [T, B] corresponding to the sampling log probability + of the chosen action w.r.t. the policy. + """ log_probs = [] for i in range(len(policy_logits)): p_shape = tf.shape(policy_logits[i]) diff --git a/rllib/agents/impala/vtrace_policy.py b/rllib/agents/impala/vtrace_tf_policy.py similarity index 99% rename from rllib/agents/impala/vtrace_policy.py rename to rllib/agents/impala/vtrace_tf_policy.py index 6480556fa..05a8e91d1 100644 --- a/rllib/agents/impala/vtrace_policy.py +++ b/rllib/agents/impala/vtrace_tf_policy.py @@ -7,7 +7,7 @@ import logging import gym import ray -from ray.rllib.agents.impala import vtrace +from ray.rllib.agents.impala import vtrace_tf as vtrace from ray.rllib.models.tf.tf_action_dist import Categorical from ray.rllib.policy.sample_batch import SampleBatch from ray.rllib.policy.tf_policy_template import build_tf_policy diff --git a/rllib/agents/impala/vtrace_torch.py b/rllib/agents/impala/vtrace_torch.py new file mode 100644 index 000000000..b74451334 --- /dev/null +++ b/rllib/agents/impala/vtrace_torch.py @@ -0,0 +1,344 @@ +# Copyright 2018 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""PyTorch version of the functions to compute V-trace off-policy actor critic +targets. + +For details and theory see: + +"IMPALA: Scalable Distributed Deep-RL with +Importance Weighted Actor-Learner Architectures" +by Espeholt, Soyer, Munos et al. + +See https://arxiv.org/abs/1802.01561 for the full paper. + +In addition to the original paper's code, changes have been made +to support MultiDiscrete action spaces. behaviour_policy_logits, +target_policy_logits and actions parameters in the entry point +multi_from_logits method accepts lists of tensors instead of just +tensors. +""" + +from ray.rllib.agents.impala.vtrace_tf import VTraceFromLogitsReturns, \ + VTraceReturns +from ray.rllib.models.torch.torch_action_dist import TorchCategorical +from ray.rllib.utils.framework import try_import_torch +from ray.rllib.utils.torch_ops import convert_to_torch_tensor + +torch, nn = try_import_torch() + + +def log_probs_from_logits_and_actions(policy_logits, + actions, + dist_class=TorchCategorical, + model=None): + return multi_log_probs_from_logits_and_actions([policy_logits], [actions], + dist_class, model)[0] + + +def multi_log_probs_from_logits_and_actions(policy_logits, actions, dist_class, + model): + """Computes action log-probs from policy logits and actions. + + In the notation used throughout documentation and comments, T refers to the + time dimension ranging from 0 to T-1. B refers to the batch size and + ACTION_SPACE refers to the list of numbers each representing a number of + actions. + + Args: + policy_logits: A list with length of ACTION_SPACE of float32 + tensors of shapes [T, B, ACTION_SPACE[0]], ..., + [T, B, ACTION_SPACE[-1]] with un-normalized log-probabilities + parameterizing a softmax policy. + actions: A list with length of ACTION_SPACE of tensors of shapes + [T, B, ...], ..., [T, B, ...] + with actions. + dist_class: Python class of the action distribution. + + Returns: + A list with length of ACTION_SPACE of float32 tensors of shapes + [T, B], ..., [T, B] corresponding to the sampling log probability + of the chosen action w.r.t. the policy. + """ + log_probs = [] + for i in range(len(policy_logits)): + p_shape = policy_logits[i].shape + a_shape = actions[i].shape + policy_logits_flat = torch.reshape(policy_logits[i], + (-1, ) + tuple(p_shape[2:])) + actions_flat = torch.reshape(actions[i], (-1, ) + tuple(a_shape[2:])) + log_probs.append( + torch.reshape( + dist_class(policy_logits_flat, model).logp(actions_flat), + a_shape[:2])) + + return log_probs + + +def from_logits(behaviour_policy_logits, + target_policy_logits, + actions, + discounts, + rewards, + values, + bootstrap_value, + dist_class=TorchCategorical, + model=None, + clip_rho_threshold=1.0, + clip_pg_rho_threshold=1.0): + """multi_from_logits wrapper used only for tests""" + + res = multi_from_logits( + [behaviour_policy_logits], [target_policy_logits], [actions], + discounts, + rewards, + values, + bootstrap_value, + dist_class, + model, + clip_rho_threshold=clip_rho_threshold, + clip_pg_rho_threshold=clip_pg_rho_threshold) + + assert len(res.behaviour_action_log_probs) == 1 + assert len(res.target_action_log_probs) == 1 + return VTraceFromLogitsReturns( + vs=res.vs, + pg_advantages=res.pg_advantages, + log_rhos=res.log_rhos, + behaviour_action_log_probs=res.behaviour_action_log_probs[0], + target_action_log_probs=res.target_action_log_probs[0], + ) + + +def multi_from_logits(behaviour_policy_logits, + target_policy_logits, + actions, + discounts, + rewards, + values, + bootstrap_value, + dist_class, + model, + behaviour_action_log_probs=None, + clip_rho_threshold=1.0, + clip_pg_rho_threshold=1.0): + """V-trace for softmax policies. + + Calculates V-trace actor critic targets for softmax polices as described in + + "IMPALA: Scalable Distributed Deep-RL with + Importance Weighted Actor-Learner Architectures" + by Espeholt, Soyer, Munos et al. + + Target policy refers to the policy we are interested in improving and + behaviour policy refers to the policy that generated the given + rewards and actions. + + In the notation used throughout documentation and comments, T refers to the + time dimension ranging from 0 to T-1. B refers to the batch size and + ACTION_SPACE refers to the list of numbers each representing a number of + actions. + + Args: + behaviour_policy_logits: A list with length of ACTION_SPACE of float32 + tensors of shapes [T, B, ACTION_SPACE[0]], ..., + [T, B, ACTION_SPACE[-1]] with un-normalized log-probabilities + parameterizing the softmax behavior policy. + target_policy_logits: A list with length of ACTION_SPACE of float32 + tensors of shapes [T, B, ACTION_SPACE[0]], ..., + [T, B, ACTION_SPACE[-1]] with un-normalized log-probabilities + parameterizing the softmax target policy. + actions: A list with length of ACTION_SPACE of tensors of shapes + [T, B, ...], ..., [T, B, ...] + with actions sampled from the behavior policy. + discounts: A float32 tensor of shape [T, B] with the discount + encountered when following the behavior policy. + rewards: A float32 tensor of shape [T, B] with the rewards generated by + following the behavior policy. + values: A float32 tensor of shape [T, B] with the value function + estimates wrt. the target policy. + bootstrap_value: A float32 of shape [B] with the value function + estimate at time T. + dist_class: action distribution class for the logits. + model: backing ModelV2 instance + behaviour_action_log_probs: Precalculated values of the behavior + actions. + clip_rho_threshold: A scalar float32 tensor with the clipping threshold + for importance weights (rho) when calculating the baseline targets + (vs). rho^bar in the paper. + clip_pg_rho_threshold: A scalar float32 tensor with the clipping + threshold on rho_s in: + \rho_s \delta log \pi(a|x) (r + \gamma v_{s+1} - V(x_s)). + + Returns: + A `VTraceFromLogitsReturns` namedtuple with the following fields: + vs: A float32 tensor of shape [T, B]. Can be used as target to train a + baseline (V(x_t) - vs_t)^2. + pg_advantages: A float 32 tensor of shape [T, B]. Can be used as an + estimate of the advantage in the calculation of policy gradients. + log_rhos: A float32 tensor of shape [T, B] containing the log + importance sampling weights (log rhos). + behaviour_action_log_probs: A float32 tensor of shape [T, B] containing + behaviour policy action log probabilities (log \mu(a_t)). + target_action_log_probs: A float32 tensor of shape [T, B] containing + target policy action probabilities (log \pi(a_t)). + """ + + behaviour_policy_logits = convert_to_torch_tensor( + behaviour_policy_logits, device="cpu") + target_policy_logits = convert_to_torch_tensor( + target_policy_logits, device="cpu") + actions = convert_to_torch_tensor(actions, device="cpu") + + for i in range(len(behaviour_policy_logits)): + # Make sure tensor ranks are as expected. + # The rest will be checked by from_action_log_probs. + assert len(behaviour_policy_logits[i].size()) == 3 + assert len(target_policy_logits[i].size()) == 3 + + target_action_log_probs = multi_log_probs_from_logits_and_actions( + target_policy_logits, actions, dist_class, model) + + if (len(behaviour_policy_logits) > 1 + or behaviour_action_log_probs is None): + # can't use precalculated values, recompute them. Note that + # recomputing won't work well for autoregressive action dists + # which may have variables not captured by 'logits' + behaviour_action_log_probs = (multi_log_probs_from_logits_and_actions( + behaviour_policy_logits, actions, dist_class, model)) + + log_rhos = get_log_rhos(target_action_log_probs, + behaviour_action_log_probs) + + vtrace_returns = from_importance_weights( + log_rhos=log_rhos, + discounts=discounts, + rewards=rewards, + values=values, + bootstrap_value=bootstrap_value, + clip_rho_threshold=clip_rho_threshold, + clip_pg_rho_threshold=clip_pg_rho_threshold) + + return VTraceFromLogitsReturns( + log_rhos=log_rhos, + behaviour_action_log_probs=behaviour_action_log_probs, + target_action_log_probs=target_action_log_probs, + **vtrace_returns._asdict()) + + +def from_importance_weights(log_rhos, + discounts, + rewards, + values, + bootstrap_value, + clip_rho_threshold=1.0, + clip_pg_rho_threshold=1.0): + """V-trace from log importance weights. + + Calculates V-trace actor critic targets as described in + + "IMPALA: Scalable Distributed Deep-RL with + Importance Weighted Actor-Learner Architectures" + by Espeholt, Soyer, Munos et al. + + In the notation used throughout documentation and comments, T refers to the + time dimension ranging from 0 to T-1. B refers to the batch size. This code + also supports the case where all tensors have the same number of additional + dimensions, e.g., `rewards` is [T, B, C], `values` is [T, B, C], + `bootstrap_value` is [B, C]. + + Args: + log_rhos: A float32 tensor of shape [T, B] representing the log + importance sampling weights, i.e. + log(target_policy(a) / behaviour_policy(a)). V-trace performs + operations on rhos in log-space for numerical stability. + discounts: A float32 tensor of shape [T, B] with discounts encountered + when following the behaviour policy. + rewards: A float32 tensor of shape [T, B] containing rewards generated + by following the behaviour policy. + values: A float32 tensor of shape [T, B] with the value function + estimates wrt. the target policy. + bootstrap_value: A float32 of shape [B] with the value function + estimate at time T. + clip_rho_threshold: A scalar float32 tensor with the clipping threshold + for importance weights (rho) when calculating the baseline targets + (vs). rho^bar in the paper. If None, no clipping is applied. + clip_pg_rho_threshold: A scalar float32 tensor with the clipping + threshold on rho_s in + \rho_s \delta log \pi(a|x) (r + \gamma v_{s+1} - V(x_s)). + If None, no clipping is applied. + + Returns: + A VTraceReturns namedtuple (vs, pg_advantages) where: + vs: A float32 tensor of shape [T, B]. Can be used as target to + train a baseline (V(x_t) - vs_t)^2. + pg_advantages: A float32 tensor of shape [T, B]. Can be used as the + advantage in the calculation of policy gradients. + """ + log_rhos = convert_to_torch_tensor(log_rhos, device="cpu") + discounts = convert_to_torch_tensor(discounts, device="cpu") + rewards = convert_to_torch_tensor(rewards, device="cpu") + values = convert_to_torch_tensor(values, device="cpu") + bootstrap_value = convert_to_torch_tensor(bootstrap_value, device="cpu") + + # Make sure tensor ranks are consistent. + rho_rank = len(log_rhos.size()) # Usually 2. + assert rho_rank == len(values.size()) + assert rho_rank - 1 == len(bootstrap_value.size()),\ + "must have rank {}".format(rho_rank - 1) + assert rho_rank == len(discounts.size()) + assert rho_rank == len(rewards.size()) + + rhos = torch.exp(log_rhos) + if clip_rho_threshold is not None: + clipped_rhos = torch.clamp_max(rhos, clip_rho_threshold) + else: + clipped_rhos = rhos + + cs = torch.clamp_max(rhos, 1.0) + # Append bootstrapped value to get [v1, ..., v_t+1] + values_t_plus_1 = torch.cat( + [values[1:], torch.unsqueeze(bootstrap_value, 0)], dim=0) + deltas = clipped_rhos * (rewards + discounts * values_t_plus_1 - values) + + vs_minus_v_xs = [torch.zeros_like(bootstrap_value)] + for i in reversed(range(len(discounts))): + discount_t, c_t, delta_t = discounts[i], cs[i], deltas[i] + vs_minus_v_xs.append(delta_t + discount_t * c_t * vs_minus_v_xs[-1]) + vs_minus_v_xs = torch.stack(vs_minus_v_xs[1:]) + # Reverse the results back to original order. + vs_minus_v_xs = torch.flip(vs_minus_v_xs, dims=[0]) + # Add V(x_s) to get v_s. + vs = vs_minus_v_xs + values + + # Advantage for policy gradient. + vs_t_plus_1 = torch.cat( + [vs[1:], torch.unsqueeze(bootstrap_value, 0)], dim=0) + if clip_pg_rho_threshold is not None: + clipped_pg_rhos = torch.clamp_max(rhos, clip_pg_rho_threshold) + else: + clipped_pg_rhos = rhos + pg_advantages = ( + clipped_pg_rhos * (rewards + discounts * vs_t_plus_1 - values)) + + # Make sure no gradients backpropagated through the returned values. + return VTraceReturns(vs=vs.detach(), pg_advantages=pg_advantages.detach()) + + +def get_log_rhos(target_action_log_probs, behaviour_action_log_probs): + """With the selected log_probs for multi-discrete actions of behavior + and target policies we compute the log_rhos for calculating the vtrace.""" + t = torch.stack(target_action_log_probs) + b = torch.stack(behaviour_action_log_probs) + log_rhos = torch.sum(t - b, dim=0) + return log_rhos diff --git a/rllib/agents/impala/vtrace_torch_policy.py b/rllib/agents/impala/vtrace_torch_policy.py new file mode 100644 index 000000000..4d1371526 --- /dev/null +++ b/rllib/agents/impala/vtrace_torch_policy.py @@ -0,0 +1,57 @@ +import logging + +from ray.rllib.utils.framework import try_import_torch + +torch, nn = try_import_torch() + +logger = logging.getLogger(__name__) + + +def make_time_major(policy, seq_lens, tensor, drop_last=False): + """Swaps batch and trajectory axis. + + Arguments: + policy: Policy reference + seq_lens: Sequence lengths if recurrent or None + tensor: A tensor or list of tensors to reshape. + drop_last: A bool indicating whether to drop the last + trajectory item. + + Returns: + res: A tensor with swapped axes or a list of tensors with + swapped axes. + """ + if isinstance(tensor, (list, tuple)): + return [ + make_time_major(policy, seq_lens, t, drop_last) for t in tensor + ] + + if policy.is_recurrent(): + B = seq_lens.shape[0] + T = tensor.shape[0] // B + else: + # Important: chop the tensor into batches at known episode cut + # boundaries. TODO(ekl) this is kind of a hack + T = policy.config["rollout_fragment_length"] + B = tensor.shape[0] // T + rs = torch.reshape(tensor, [B, T] + list(tensor.shape[1:])) + + # Swap B and T axes. + res = torch.transpose(rs, 1, 0) + + if drop_last: + return res[:-1] + return res + + +def choose_optimizer(policy, config): + if policy.config["opt_type"] == "adam": + return torch.optim.Adam( + params=policy.model.parameters(), lr=policy.cur_lr) + else: + return torch.optim.RMSProp( + params=policy.model.parameters(), + lr=policy.cur_lr, + weight_decay=config["decay"], + momentum=config["momentum"], + eps=config["epsilon"]) diff --git a/rllib/agents/ppo/appo.py b/rllib/agents/ppo/appo.py index 93d485cb2..535a2013a 100644 --- a/rllib/agents/ppo/appo.py +++ b/rllib/agents/ppo/appo.py @@ -1,4 +1,4 @@ -from ray.rllib.agents.ppo.appo_policy import AsyncPPOTFPolicy +from ray.rllib.agents.ppo.appo_tf_policy import AsyncPPOTFPolicy from ray.rllib.agents.trainer import with_base_config from ray.rllib.agents.ppo.ppo import update_kl from ray.rllib.agents import impala @@ -81,10 +81,24 @@ def initialize_target(trainer): * trainer.config["minibatch_buffer_size"] +def get_policy_class(config): + if config.get("use_pytorch") is True: + from ray.rllib.agents.ppo.appo_torch_policy import AsyncPPOTorchPolicy + return AsyncPPOTorchPolicy + else: + return AsyncPPOTFPolicy + + +def validate_config(config): + if config["entropy_coeff"] < 0: + raise ValueError("`entropy_coeff` must be >= 0.0!") + + APPOTrainer = impala.ImpalaTrainer.with_updates( name="APPO", default_config=DEFAULT_CONFIG, + validate_config=validate_config, default_policy=AsyncPPOTFPolicy, - get_policy_class=lambda _: AsyncPPOTFPolicy, + get_policy_class=get_policy_class, after_init=initialize_target, after_optimizer_step=update_target_and_kl) diff --git a/rllib/agents/ppo/appo_policy.py b/rllib/agents/ppo/appo_tf_policy.py similarity index 97% rename from rllib/agents/ppo/appo_policy.py rename to rllib/agents/ppo/appo_tf_policy.py index 0662bf0f0..1cddbc6fa 100644 --- a/rllib/agents/ppo/appo_policy.py +++ b/rllib/agents/ppo/appo_tf_policy.py @@ -6,8 +6,8 @@ import numpy as np import logging import gym -from ray.rllib.agents.impala import vtrace -from ray.rllib.agents.impala.vtrace_policy import _make_time_major, \ +from ray.rllib.agents.impala import vtrace_tf as vtrace +from ray.rllib.agents.impala.vtrace_tf_policy import _make_time_major, \ clip_gradients, validate_config, choose_optimizer from ray.rllib.evaluation.postprocessing import Postprocessing from ray.rllib.models.tf.tf_action_dist import Categorical @@ -209,7 +209,8 @@ def build_appo_model(policy, obs_space, action_space, config): logit_dim, config["model"], name=POLICY_SCOPE, - framework="tf") + framework="torch" if config["use_pytorch"] else "tf") + policy.model_variables = policy.model.variables() policy.target_model = ModelCatalog.get_model_v2( obs_space, @@ -217,7 +218,8 @@ def build_appo_model(policy, obs_space, action_space, config): logit_dim, config["model"], name=TARGET_POLICY_SCOPE, - framework="tf") + framework="torch" if config["use_pytorch"] else "tf") + policy.target_model_variables = policy.target_model.variables() return policy.model diff --git a/rllib/agents/ppo/appo_torch_policy.py b/rllib/agents/ppo/appo_torch_policy.py new file mode 100644 index 000000000..0dfb3f915 --- /dev/null +++ b/rllib/agents/ppo/appo_torch_policy.py @@ -0,0 +1,411 @@ +"""Adapted from VTraceTFPolicy to use the PPO surrogate loss. + +Keep in sync with changes to VTraceTFPolicy.""" + +import numpy as np +import logging +import gym + +from ray.rllib.agents.a3c.a3c_torch_policy import apply_grad_clipping +import ray.rllib.agents.impala.vtrace_torch as vtrace +from ray.rllib.agents.impala.vtrace_torch_policy import make_time_major, \ + choose_optimizer +from ray.rllib.agents.ppo.appo_tf_policy import build_appo_model, \ + postprocess_trajectory +from ray.rllib.agents.ppo.ppo_torch_policy import ValueNetworkMixin, \ + KLCoeffMixin +from ray.rllib.evaluation.postprocessing import Postprocessing +from ray.rllib.models.torch.torch_action_dist import TorchCategorical +from ray.rllib.policy.sample_batch import SampleBatch +from ray.rllib.policy.torch_policy import LearningRateSchedule +from ray.rllib.policy.torch_policy_template import build_torch_policy +from ray.rllib.utils.explained_variance import explained_variance +from ray.rllib.utils.framework import try_import_torch +from ray.rllib.utils.torch_ops import global_norm, sequence_mask + +torch, nn = try_import_torch() + +logger = logging.getLogger(__name__) + + +class PPOSurrogateLoss: + """Loss used when V-trace is disabled. + + Arguments: + prev_actions_logp: A float32 tensor of shape [T, B]. + actions_logp: A float32 tensor of shape [T, B]. + action_kl: A float32 tensor of shape [T, B]. + actions_entropy: A float32 tensor of shape [T, B]. + values: A float32 tensor of shape [T, B]. + valid_mask: A bool tensor of valid RNN input elements (#2992). + advantages: A float32 tensor of shape [T, B]. + value_targets: A float32 tensor of shape [T, B]. + vf_loss_coeff (float): Coefficient of the value function loss. + entropy_coeff (float): Coefficient of the entropy regularizer. + clip_param (float): Clip parameter. + cur_kl_coeff (float): Coefficient for KL loss. + use_kl_loss (bool): If true, use KL loss. + """ + + def __init__(self, + prev_actions_logp, + actions_logp, + action_kl, + actions_entropy, + values, + valid_mask, + advantages, + value_targets, + vf_loss_coeff=0.5, + entropy_coeff=0.01, + clip_param=0.3, + cur_kl_coeff=None, + use_kl_loss=False): + + if valid_mask is not None: + num_valid = torch.sum(valid_mask) + + def reduce_mean_valid(t): + return torch.sum(t * valid_mask) / num_valid + + else: + + def reduce_mean_valid(t): + return torch.mean(t) + + logp_ratio = torch.exp(actions_logp - prev_actions_logp) + + surrogate_loss = torch.min( + advantages * logp_ratio, + advantages * torch.clamp(logp_ratio, 1 - clip_param, + 1 + clip_param)) + + self.mean_kl = reduce_mean_valid(action_kl) + self.pi_loss = -reduce_mean_valid(surrogate_loss) + + # The baseline loss + delta = values - value_targets + self.value_targets = value_targets + self.vf_loss = 0.5 * reduce_mean_valid(torch.pow(delta, 2.0)) + + # The entropy loss + self.entropy = reduce_mean_valid(actions_entropy) + + # The summed weighted loss + self.total_loss = (self.pi_loss + self.vf_loss * vf_loss_coeff - + self.entropy * entropy_coeff) + + # Optional additional KL Loss + if use_kl_loss: + self.total_loss += cur_kl_coeff * self.mean_kl + + +class VTraceSurrogateLoss: + def __init__(self, + actions, + prev_actions_logp, + actions_logp, + old_policy_actions_logp, + action_kl, + actions_entropy, + dones, + behaviour_logits, + old_policy_behaviour_logits, + target_logits, + discount, + rewards, + values, + bootstrap_value, + dist_class, + model, + valid_mask, + vf_loss_coeff=0.5, + entropy_coeff=0.01, + clip_rho_threshold=1.0, + clip_pg_rho_threshold=1.0, + clip_param=0.3, + cur_kl_coeff=None, + use_kl_loss=False): + """APPO Loss, with IS modifications and V-trace for Advantage Estimation + + VTraceLoss takes tensors of shape [T, B, ...], where `B` is the + batch_size. The reason we need to know `B` is for V-trace to properly + handle episode cut boundaries. + + Arguments: + actions: An int|float32 tensor of shape [T, B, logit_dim]. + prev_actions_logp: A float32 tensor of shape [T, B]. + actions_logp: A float32 tensor of shape [T, B]. + old_policy_actions_logp: A float32 tensor of shape [T, B]. + action_kl: A float32 tensor of shape [T, B]. + actions_entropy: A float32 tensor of shape [T, B]. + dones: A bool tensor of shape [T, B]. + behaviour_logits: A float32 tensor of shape [T, B, logit_dim]. + old_policy_behaviour_logits: A float32 tensor of shape + [T, B, logit_dim]. + target_logits: A float32 tensor of shape [T, B, logit_dim]. + discount: A float32 scalar. + rewards: A float32 tensor of shape [T, B]. + values: A float32 tensor of shape [T, B]. + bootstrap_value: A float32 tensor of shape [B]. + dist_class: action distribution class for logits. + model: backing ModelV2 instance + valid_mask: A bool tensor of valid RNN input elements (#2992). + vf_loss_coeff (float): Coefficient of the value function loss. + entropy_coeff (float): Coefficient of the entropy regularizer. + clip_param (float): Clip parameter. + cur_kl_coeff (float): Coefficient for KL loss. + use_kl_loss (bool): If true, use KL loss. + """ + + if valid_mask is not None: + num_valid = torch.sum(valid_mask) + + def reduce_mean_valid(t): + return torch.sum(t * valid_mask) / num_valid + + else: + + def reduce_mean_valid(t): + return torch.mean(t) + + # Compute vtrace on the CPU for better perf. + self.vtrace_returns = vtrace.multi_from_logits( + behaviour_policy_logits=behaviour_logits, + target_policy_logits=old_policy_behaviour_logits, + actions=torch.unbind(actions, dim=2), + discounts=(1.0 - dones.float()) * discount, + rewards=rewards, + values=values, + bootstrap_value=bootstrap_value, + dist_class=dist_class, + model=model, + clip_rho_threshold=clip_rho_threshold, + clip_pg_rho_threshold=clip_pg_rho_threshold) + + self.is_ratio = torch.clamp( + torch.exp(prev_actions_logp - old_policy_actions_logp), 0.0, 2.0) + logp_ratio = self.is_ratio * torch.exp(actions_logp - + prev_actions_logp) + + advantages = self.vtrace_returns.pg_advantages + surrogate_loss = torch.min( + advantages * logp_ratio, + advantages * torch.clamp(logp_ratio, 1 - clip_param, + 1 + clip_param)) + + self.mean_kl = reduce_mean_valid(action_kl) + self.pi_loss = -reduce_mean_valid(surrogate_loss) + + # The baseline loss + delta = values - self.vtrace_returns.vs + self.value_targets = self.vtrace_returns.vs + self.vf_loss = 0.5 * reduce_mean_valid(torch.pow(delta, 2.0)) + + # The entropy loss + self.entropy = reduce_mean_valid(actions_entropy) + + # The summed weighted loss + self.total_loss = (self.pi_loss + self.vf_loss * vf_loss_coeff - + self.entropy * entropy_coeff) + + # Optional additional KL Loss + if use_kl_loss: + self.total_loss += cur_kl_coeff * self.mean_kl + + +def build_appo_surrogate_loss(policy, model, dist_class, train_batch): + model_out, _ = model.from_batch(train_batch) + action_dist = dist_class(model_out, model) + + if isinstance(policy.action_space, gym.spaces.Discrete): + is_multidiscrete = False + output_hidden_shape = [policy.action_space.n] + elif isinstance(policy.action_space, + gym.spaces.multi_discrete.MultiDiscrete): + is_multidiscrete = True + output_hidden_shape = policy.action_space.nvec.astype(np.int32) + else: + is_multidiscrete = False + output_hidden_shape = 1 + + def _make_time_major(*args, **kw): + return make_time_major(policy, train_batch.get("seq_lens"), *args, + **kw) + + actions = train_batch[SampleBatch.ACTIONS] + dones = train_batch[SampleBatch.DONES] + rewards = train_batch[SampleBatch.REWARDS] + behaviour_logits = train_batch[SampleBatch.ACTION_DIST_INPUTS] + + target_model_out, _ = policy.target_model.from_batch(train_batch) + old_policy_behaviour_logits = target_model_out.detach() + + unpacked_behaviour_logits = torch.split( + behaviour_logits, output_hidden_shape, dim=1) + unpacked_old_policy_behaviour_logits = torch.split( + old_policy_behaviour_logits, output_hidden_shape, dim=1) + unpacked_outputs = torch.split(model_out, output_hidden_shape, dim=1) + old_policy_action_dist = dist_class(old_policy_behaviour_logits, model) + prev_action_dist = dist_class(behaviour_logits, policy.model) + values = policy.model.value_function() + + policy.model_vars = policy.model.variables() + policy.target_model_vars = policy.target_model.variables() + + if policy.is_recurrent(): + max_seq_len = torch.max(train_batch["seq_lens"]) - 1 + mask = sequence_mask(train_batch["seq_lens"], max_seq_len) + mask = torch.reshape(mask, [-1]) + else: + mask = torch.ones_like(rewards) + + if policy.config["vtrace"]: + logger.debug("Using V-Trace surrogate loss (vtrace=True)") + + # Prepare actions for loss + loss_actions = actions if is_multidiscrete else torch.unsqueeze( + actions, dim=1) + + # Prepare KL for Loss + mean_kl = _make_time_major( + old_policy_action_dist.multi_kl(action_dist), drop_last=True) + + policy.loss = VTraceSurrogateLoss( + actions=_make_time_major(loss_actions, drop_last=True), + prev_actions_logp=_make_time_major( + prev_action_dist.logp(actions), drop_last=True), + actions_logp=_make_time_major( + action_dist.logp(actions), drop_last=True), + old_policy_actions_logp=_make_time_major( + old_policy_action_dist.logp(actions), drop_last=True), + action_kl=torch.mean(mean_kl, dim=0) + if is_multidiscrete else mean_kl, + actions_entropy=_make_time_major( + action_dist.multi_entropy(), drop_last=True), + dones=_make_time_major(dones, drop_last=True), + behaviour_logits=_make_time_major( + unpacked_behaviour_logits, drop_last=True), + old_policy_behaviour_logits=_make_time_major( + unpacked_old_policy_behaviour_logits, drop_last=True), + target_logits=_make_time_major(unpacked_outputs, drop_last=True), + discount=policy.config["gamma"], + rewards=_make_time_major(rewards, drop_last=True), + values=_make_time_major(values, drop_last=True), + bootstrap_value=_make_time_major(values)[-1], + dist_class=TorchCategorical if is_multidiscrete else dist_class, + model=policy.model, + valid_mask=_make_time_major(mask, drop_last=True), + vf_loss_coeff=policy.config["vf_loss_coeff"], + entropy_coeff=policy.config["entropy_coeff"], + clip_rho_threshold=policy.config["vtrace_clip_rho_threshold"], + clip_pg_rho_threshold=policy.config[ + "vtrace_clip_pg_rho_threshold"], + clip_param=policy.config["clip_param"], + cur_kl_coeff=policy.kl_coeff, + use_kl_loss=policy.config["use_kl_loss"]) + else: + logger.debug("Using PPO surrogate loss (vtrace=False)") + + # Prepare KL for Loss + mean_kl = _make_time_major(prev_action_dist.multi_kl(action_dist)) + + policy.loss = PPOSurrogateLoss( + prev_actions_logp=_make_time_major(prev_action_dist.logp(actions)), + actions_logp=_make_time_major(action_dist.logp(actions)), + action_kl=torch.mean(mean_kl, dim=0) + if is_multidiscrete else mean_kl, + actions_entropy=_make_time_major(action_dist.multi_entropy()), + values=_make_time_major(values), + valid_mask=_make_time_major(mask), + advantages=_make_time_major( + train_batch[Postprocessing.ADVANTAGES]), + value_targets=_make_time_major( + train_batch[Postprocessing.VALUE_TARGETS]), + vf_loss_coeff=policy.config["vf_loss_coeff"], + entropy_coeff=policy.config["entropy_coeff"], + clip_param=policy.config["clip_param"], + cur_kl_coeff=policy.kl_coeff, + use_kl_loss=policy.config["use_kl_loss"]) + + return policy.loss.total_loss + + +def stats(policy, train_batch): + values_batched = make_time_major( + policy, + train_batch.get("seq_lens"), + policy.model.value_function(), + drop_last=policy.config["vtrace"]) + + stats_dict = { + "cur_lr": policy.cur_lr, + "policy_loss": policy.loss.pi_loss, + "entropy": policy.loss.entropy, + "var_gnorm": global_norm(policy.model.trainable_variables()), + "vf_loss": policy.loss.vf_loss, + "vf_explained_var": explained_variance( + torch.reshape(policy.loss.value_targets, [-1]), + torch.reshape(values_batched, [-1]), + framework="torch"), + } + + if policy.config["vtrace"]: + is_stat_mean = torch.mean(policy.loss.is_ratio, [0, 1]) + is_stat_var = torch.var(policy.loss.is_ratio, [0, 1]) + stats_dict.update({"mean_IS": is_stat_mean}) + stats_dict.update({"var_IS": is_stat_var}) + + if policy.config["use_kl_loss"]: + stats_dict.update({"kl": policy.loss.mean_kl}) + stats_dict.update({"KL_Coeff": policy.kl_coeff}) + + return stats_dict + + +class TargetNetworkMixin: + def __init__(self, obs_space, action_space, config): + def do_update(): + # Update_target_fn will be called periodically to copy Q network to + # target Q network. + assert len(self.model_variables) == \ + len(self.target_model_variables), \ + (self.model_variables, self.target_model_variables) + self.target_model.load_state_dict(self.model.state_dict()) + + self.update_target = do_update + + +def add_values(policy, input_dict, state_batches, model, action_dist): + out = {} + if not policy.config["vtrace"]: + out[SampleBatch.VF_PREDS] = policy.model.value_function() + return out + + +def setup_early_mixins(policy, obs_space, action_space, config): + LearningRateSchedule.__init__(policy, config["lr"], config["lr_schedule"]) + + +def setup_late_mixins(policy, obs_space, action_space, config): + KLCoeffMixin.__init__(policy, config) + ValueNetworkMixin.__init__(policy, obs_space, action_space, config) + TargetNetworkMixin.__init__(policy, obs_space, action_space, config) + + +AsyncPPOTorchPolicy = build_torch_policy( + name="AsyncPPOTorchPolicy", + loss_fn=build_appo_surrogate_loss, + stats_fn=stats, + postprocess_fn=postprocess_trajectory, + extra_action_out_fn=add_values, + extra_grad_process_fn=apply_grad_clipping, + optimizer_fn=choose_optimizer, + before_init=setup_early_mixins, + after_init=setup_late_mixins, + make_model=build_appo_model, + mixins=[ + LearningRateSchedule, KLCoeffMixin, TargetNetworkMixin, + ValueNetworkMixin + ], + get_batch_divisibility_req=lambda p: p.config["rollout_fragment_length"]) diff --git a/rllib/agents/ppo/tests/test_appo.py b/rllib/agents/ppo/tests/test_appo.py new file mode 100644 index 000000000..0129802c2 --- /dev/null +++ b/rllib/agents/ppo/tests/test_appo.py @@ -0,0 +1,42 @@ +import unittest + +import ray +import ray.rllib.agents.ppo as ppo +from ray.rllib.utils.framework import try_import_tf +from ray.rllib.utils.test_utils import framework_iterator + +tf = try_import_tf() + + +class TestAPPO(unittest.TestCase): + @classmethod + def setUpClass(cls): + ray.init() + + @classmethod + def tearDownClass(cls): + ray.shutdown() + + def test_appo_compilation(self): + """Test whether an APPOTrainer can be built with both frameworks.""" + config = ppo.appo.DEFAULT_CONFIG.copy() + config["num_workers"] = 1 + num_iterations = 2 + + for _ in framework_iterator(config, frameworks=("torch", "tf")): + _config = config.copy() + trainer = ppo.APPOTrainer(config=_config, env="CartPole-v0") + for i in range(num_iterations): + print(trainer.train()) + + _config = config.copy() + _config["vtrace"] = True + trainer = ppo.APPOTrainer(config=_config, env="CartPole-v0") + for i in range(num_iterations): + print(trainer.train()) + + +if __name__ == "__main__": + import pytest + import sys + sys.exit(pytest.main(["-v", __file__])) diff --git a/rllib/policy/torch_policy_template.py b/rllib/policy/torch_policy_template.py index 8435f3894..542214efc 100644 --- a/rllib/policy/torch_policy_template.py +++ b/rllib/policy/torch_policy_template.py @@ -24,6 +24,7 @@ def build_torch_policy(name, after_init=None, action_sampler_fn=None, action_distribution_fn=None, + make_model=None, make_model_and_action_dist=None, apply_gradients_fn=None, mixins=None, @@ -62,11 +63,16 @@ def build_torch_policy(name, a) distribution inputs (parameters), b) a dist-class to generate an action distribution object from, and c) internal-state outputs (empty list if not applicable). + make_model (Optional[callable]): Optional func that + takes the same arguments as Policy.__init__ and returns a model + instance. The distribution class will be determined automatically. + Note: Only one of `make_model` or `make_model_and_action_dist` + should be provided. make_model_and_action_dist (Optional[callable]): Optional func that takes the same arguments as Policy.__init__ and returns a tuple - of model instance and torch action distribution class. If not - specified, the default model and action dist from the catalog will - be used. + of model instance and torch action distribution class. + Note: Only one of `make_model` or `make_model_and_action_dist` + should be provided. apply_gradients_fn (Optional[callable]): Optional callable that takes a grads list and applies these to the Model's parameters. mixins (list): list of any class mixins for the returned policy class. @@ -91,13 +97,17 @@ def build_torch_policy(name, if before_init: before_init(self, obs_space, action_space, config) - if make_model_and_action_dist: + # Model is customized (use default action dist class). + if make_model: + assert make_model_and_action_dist is None + self.model = make_model(self, obs_space, action_space, config) + dist_class, _ = ModelCatalog.get_action_dist( + action_space, self.config["model"], framework="torch") + # Model and action dist class are customized. + elif make_model_and_action_dist: self.model, dist_class = make_model_and_action_dist( self, obs_space, action_space, config) - # Make sure, we passed in a correct Model factory. - assert isinstance(self.model, TorchModelV2), \ - "ERROR: TorchPolicy::make_model_and_action_dist must " \ - "return a TorchModelV2 object!" + # Use default model and default action dist. else: dist_class, logit_dim = ModelCatalog.get_action_dist( action_space, self.config["model"], framework="torch") @@ -109,6 +119,10 @@ def build_torch_policy(name, framework="torch", **self.config["model"].get("custom_options", {})) + # Make sure, we passed in a correct Model factory. + assert isinstance(self.model, TorchModelV2), \ + "ERROR: Generated Model must be a TorchModelV2 object!" + TorchPolicy.__init__( self, observation_space=obs_space, diff --git a/rllib/tuned_examples/regression_tests/pendulum-appo-vtrace.yaml b/rllib/tuned_examples/pendulum-appo-tf.yaml similarity index 61% rename from rllib/tuned_examples/regression_tests/pendulum-appo-vtrace.yaml rename to rllib/tuned_examples/pendulum-appo-tf.yaml index ae7f5deaa..29f6d505a 100644 --- a/rllib/tuned_examples/regression_tests/pendulum-appo-vtrace.yaml +++ b/rllib/tuned_examples/pendulum-appo-tf.yaml @@ -1,11 +1,12 @@ -pendulum-appo-vt: +pendulum-appo-vtrace-tf: env: Pendulum-v0 run: APPO stop: - episode_reward_mean: -1200 # just check it learns a bit + episode_reward_mean: -1000 # just check it learns a bit timesteps_total: 500000 config: - vtrace: False + use_pytorch: false + vtrace: true num_gpus: 0 num_workers: 1 lambda: 0.1 @@ -15,6 +16,6 @@ pendulum-appo-vt: minibatch_buffer_size: 16 num_sgd_iter: 10 model: - fcnet_hiddens: [64, 64] - batch_mode: complete_episodes + fcnet_hiddens: [256, 256] + batch_mode: truncate_episodes observation_filter: MeanStdFilter diff --git a/rllib/tuned_examples/pendulum-appo-torch.yaml b/rllib/tuned_examples/pendulum-appo-torch.yaml new file mode 100644 index 000000000..73f4b6348 --- /dev/null +++ b/rllib/tuned_examples/pendulum-appo-torch.yaml @@ -0,0 +1,21 @@ +pendulum-appo-vtrace-torch: + env: Pendulum-v0 + run: APPO + stop: + episode_reward_mean: -1000 # just check it learns a bit + timesteps_total: 500000 + config: + use_pytorch: true + vtrace: true + num_gpus: 0 + num_workers: 1 + lambda: 0.1 + gamma: 0.95 + lr: 0.0003 + train_batch_size: 100 + minibatch_buffer_size: 16 + num_sgd_iter: 10 + model: + fcnet_hiddens: [256, 256] + batch_mode: truncate_episodes + observation_filter: MeanStdFilter diff --git a/rllib/tuned_examples/regression_tests/cartpole-appo.yaml b/rllib/tuned_examples/regression_tests/cartpole-appo-tf.yaml similarity index 66% rename from rllib/tuned_examples/regression_tests/cartpole-appo.yaml rename to rllib/tuned_examples/regression_tests/cartpole-appo-tf.yaml index f7bf77293..183dcb05d 100644 --- a/rllib/tuned_examples/regression_tests/cartpole-appo.yaml +++ b/rllib/tuned_examples/regression_tests/cartpole-appo-tf.yaml @@ -1,10 +1,11 @@ -cartpole-appo: +cartpole-appo-tf: env: CartPole-v0 run: APPO stop: - episode_reward_mean: 100 - timesteps_total: 100000 + episode_reward_mean: 150 + timesteps_total: 200000 config: + use_pytorch: false rollout_fragment_length: 10 train_batch_size: 10 num_envs_per_worker: 5 diff --git a/rllib/tuned_examples/regression_tests/cartpole-appo-torch.yaml b/rllib/tuned_examples/regression_tests/cartpole-appo-torch.yaml new file mode 100644 index 000000000..a1db540a7 --- /dev/null +++ b/rllib/tuned_examples/regression_tests/cartpole-appo-torch.yaml @@ -0,0 +1,14 @@ +cartpole-appo-torch: + env: CartPole-v0 + run: APPO + stop: + episode_reward_mean: 150 + timesteps_total: 200000 + config: + use_pytorch: true + rollout_fragment_length: 10 + train_batch_size: 10 + num_envs_per_worker: 5 + num_workers: 1 + num_gpus: 0 + vtrace: false diff --git a/rllib/tuned_examples/regression_tests/cartpole-appo-vtrace.yaml b/rllib/tuned_examples/regression_tests/cartpole-appo-vtrace-tf.yaml similarity index 64% rename from rllib/tuned_examples/regression_tests/cartpole-appo-vtrace.yaml rename to rllib/tuned_examples/regression_tests/cartpole-appo-vtrace-tf.yaml index 12ac84dd0..023dc56a4 100644 --- a/rllib/tuned_examples/regression_tests/cartpole-appo-vtrace.yaml +++ b/rllib/tuned_examples/regression_tests/cartpole-appo-vtrace-tf.yaml @@ -1,10 +1,11 @@ -cartpole-appo-vt: +cartpole-appo-vtrace-tf: env: CartPole-v0 run: APPO stop: - episode_reward_mean: 100 - timesteps_total: 100000 + episode_reward_mean: 150 + timesteps_total: 200000 config: + use_pytorch: false rollout_fragment_length: 10 train_batch_size: 10 num_envs_per_worker: 5 diff --git a/rllib/tuned_examples/regression_tests/cartpole-appo-vtrace-torch.yaml b/rllib/tuned_examples/regression_tests/cartpole-appo-vtrace-torch.yaml new file mode 100644 index 000000000..4466621e0 --- /dev/null +++ b/rllib/tuned_examples/regression_tests/cartpole-appo-vtrace-torch.yaml @@ -0,0 +1,14 @@ +cartpole-appo-vtrace-torch: + env: CartPole-v0 + run: APPO + stop: + episode_reward_mean: 150 + timesteps_total: 200000 + config: + use_pytorch: true + rollout_fragment_length: 10 + train_batch_size: 10 + num_envs_per_worker: 5 + num_workers: 1 + num_gpus: 0 + vtrace: true diff --git a/rllib/utils/torch_ops.py b/rllib/utils/torch_ops.py index 269fc2ac0..08452f426 100644 --- a/rllib/utils/torch_ops.py +++ b/rllib/utils/torch_ops.py @@ -14,6 +14,24 @@ except (ImportError, ModuleNotFoundError) as e: raise e +def global_norm(tensors): + """Returns the global L2 norm over a list of tensors. + + output = sqrt(SUM(t ** 2 for t in tensors)), + where SUM reduces over all tensors and over all elements in tensors. + + Args: + tensors (List[torch.Tensor]): The list of tensors to calculate the + global norm over. + """ + # List of single tensors' L2 norms: SQRT(SUM(xi^2)) over all xi in tensor. + single_l2s = [ + torch.pow(torch.sum(torch.pow(t, 2.0)), 0.5) for t in tensors + ] + # Compute global norm from all single tensors' L2 norms. + return torch.pow(sum(torch.pow(l2, 2.0) for l2 in single_l2s), 0.5) + + def huber_loss(x, delta=1.0): """Reference: https://en.wikipedia.org/wiki/Huber_loss""" return torch.where(