From d6cb7d865e010619536f7911d57ae260c7d0813e Mon Sep 17 00:00:00 2001 From: Sven Mika Date: Mon, 20 Apr 2020 10:03:25 +0200 Subject: [PATCH] [RLlib] Torch DQN (APEX) TD-Error/prio. replay fixes. (#8082) PyTorch APEX_DQN with Prioritized Replay enabled would not work properly due to the td_error not being retrievable by the AsyncReplayOptimizer. --- rllib/agents/dqn/dqn_torch_policy.py | 10 +++++++++- rllib/agents/dqn/tests/test_apex.py | 16 +++++++++++++--- rllib/optimizers/async_replay_optimizer.py | 6 +++++- rllib/tuned_examples/atari-apex.yaml | 2 +- 4 files changed, 28 insertions(+), 6 deletions(-) diff --git a/rllib/agents/dqn/dqn_torch_policy.py b/rllib/agents/dqn/dqn_torch_policy.py index 606da13a8..7050b3cfa 100644 --- a/rllib/agents/dqn/dqn_torch_policy.py +++ b/rllib/agents/dqn/dqn_torch_policy.py @@ -244,6 +244,14 @@ def compute_q_values(policy, model, obs, explore, is_training=False): return q_values +def grad_process_and_td_error_fn(policy, optimizer, loss): + # Clip grads if configured. + info = apply_grad_clipping(policy, optimizer, loss) + # Add td-error to info dict. + info["td_error"] = policy.q_loss.td_error + return info + + def extra_action_out_fn(policy, input_dict, state_batches, model, action_dist): return {"q_values": policy.q_values} @@ -257,7 +265,7 @@ DQNTorchPolicy = build_torch_policy( stats_fn=build_q_stats, postprocess_fn=postprocess_nstep_and_prio, optimizer_fn=adam_optimizer, - extra_grad_process_fn=apply_grad_clipping, + extra_grad_process_fn=grad_process_and_td_error_fn, extra_action_out_fn=extra_action_out_fn, before_init=setup_early_mixins, after_init=after_init, diff --git a/rllib/agents/dqn/tests/test_apex.py b/rllib/agents/dqn/tests/test_apex.py index 9ae6cdca9..b5fc2d89f 100644 --- a/rllib/agents/dqn/tests/test_apex.py +++ b/rllib/agents/dqn/tests/test_apex.py @@ -14,19 +14,29 @@ class TestApex(unittest.TestCase): def tearDown(self): ray.shutdown() - def test_apex_epsilon_distribution(self): + def test_apex_compilation_and_per_worker_epsilon_values(self): + """Test whether an APEX-DQNTrainer can be built on all frameworks.""" config = apex.APEX_DEFAULT_CONFIG.copy() config["num_workers"] = 3 + config["prioritized_replay"] = True config["optimizer"]["num_replay_buffer_shards"] = 1 + num_iterations = 1 - for _ in framework_iterator(config): - trainer = apex.ApexTrainer(config, env="CartPole-v0") + for _ in framework_iterator(config, ("torch", "tf", "eager")): + plain_config = config.copy() + trainer = apex.ApexTrainer(config=plain_config, env="CartPole-v0") + + # Test per-worker epsilon distribution. infos = trainer.workers.foreach_policy( lambda p, _: p.get_exploration_info()) eps = [i["cur_epsilon"] for i in infos] assert np.allclose(eps, [1.0, 0.016190862, 0.00065536, 2.6527108e-05]) + for i in range(num_iterations): + results = trainer.train() + print(results) + if __name__ == "__main__": import sys diff --git a/rllib/optimizers/async_replay_optimizer.py b/rllib/optimizers/async_replay_optimizer.py index e2fa47042..8a87ce456 100644 --- a/rllib/optimizers/async_replay_optimizer.py +++ b/rllib/optimizers/async_replay_optimizer.py @@ -15,6 +15,7 @@ import ray from ray.exceptions import RayError from ray.util.iter import ParallelIteratorWorker from ray.rllib.evaluation.metrics import get_learner_stats +from ray.rllib.policy.policy import LEARNER_STATS_KEY from ray.rllib.policy.sample_batch import SampleBatch, DEFAULT_POLICY_ID, \ MultiAgentBatch from ray.rllib.optimizers.policy_optimizer import PolicyOptimizer @@ -462,8 +463,11 @@ class LearnerThread(threading.Thread): with self.grad_timer: grad_out = self.local_worker.learn_on_batch(replay) for pid, info in grad_out.items(): + td_error = info.get( + "td_error", + info[LEARNER_STATS_KEY].get("td_error")) prio_dict[pid] = (replay.policy_batches[pid].data.get( - "batch_indexes"), info.get("td_error")) + "batch_indexes"), td_error) self.stats[pid] = get_learner_stats(info) self.grad_timer.push_units_processed(replay.count) self.outqueue.put((ra, prio_dict, replay.count)) diff --git a/rllib/tuned_examples/atari-apex.yaml b/rllib/tuned_examples/atari-apex.yaml index 73c3e8278..779b677b0 100644 --- a/rllib/tuned_examples/atari-apex.yaml +++ b/rllib/tuned_examples/atari-apex.yaml @@ -22,7 +22,7 @@ apex: epsilon_timesteps: 200000 prioritized_replay_alpha: 0.5 final_prioritized_replay_beta: 1.0 - final_prioritized_replay_beta_annealing_timesteps: 2000000 + prioritized_replay_beta_annealing_timesteps: 2000000 num_gpus: 1