[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.
This commit is contained in:
Sven Mika
2020-04-20 10:03:25 +02:00
committed by GitHub
parent c8b9a357f2
commit d6cb7d865e
4 changed files with 28 additions and 6 deletions
+9 -1
View File
@@ -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,
+13 -3
View File
@@ -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
+5 -1
View File
@@ -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))
+1 -1
View File
@@ -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