mirror of
https://github.com/wassname/ray.git
synced 2026-08-09 12:20:09 +08:00
[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:
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user