diff --git a/python/ray/rllib/agents/dqn/dqn.py b/python/ray/rllib/agents/dqn/dqn.py index cdace0865..db5d261ee 100644 --- a/python/ray/rllib/agents/dqn/dqn.py +++ b/python/ray/rllib/agents/dqn/dqn.py @@ -7,7 +7,6 @@ import time from ray.rllib import optimizers from ray.rllib.agents.agent import Agent, with_common_config from ray.rllib.agents.dqn.dqn_policy_graph import DQNPolicyGraph -from ray.rllib.evaluation.metrics import collect_metrics from ray.rllib.utils.schedules import ConstantSchedule, LinearSchedule OPTIMIZER_SHARED_CONFIGS = [ @@ -220,14 +219,12 @@ class DQNAgent(Agent): if self.config["per_worker_exploration"]: # Only collect metrics from the third of workers with lowest eps - result = collect_metrics( - self.local_evaluator, - self.remote_evaluators[-len(self.remote_evaluators) // 3:], - timeout_seconds=self.config["collect_metrics_timeout"]) + result = self.optimizer.collect_metrics( + timeout_seconds=self.config["collect_metrics_timeout"], + selected_evaluators=self.remote_evaluators[ + -len(self.remote_evaluators) // 3:]) else: - result = collect_metrics( - self.local_evaluator, - self.remote_evaluators, + result = self.optimizer.collect_metrics( timeout_seconds=self.config["collect_metrics_timeout"]) result.update( diff --git a/python/ray/rllib/optimizers/policy_optimizer.py b/python/ray/rllib/optimizers/policy_optimizer.py index 3f958c4d6..6c2a5787d 100644 --- a/python/ray/rllib/optimizers/policy_optimizer.py +++ b/python/ray/rllib/optimizers/policy_optimizer.py @@ -86,11 +86,18 @@ class PolicyOptimizer(object): "num_steps_sampled": self.num_steps_sampled, } - def collect_metrics(self, timeout_seconds, min_history=100): + def collect_metrics(self, + timeout_seconds, + min_history=100, + selected_evaluators=None): """Returns evaluator and optimizer stats. Arguments: + timeout_seconds (int): Max wait time for a evaluator before + dropping its results. This usually indicates a hung evaluator. min_history (int): Min history length to smooth results over. + selected_evaluators (list): Override the list of remote evaluators + to collect metrics from. Returns: res (dict): A training result dict from evaluator metrics with @@ -98,7 +105,7 @@ class PolicyOptimizer(object): """ episodes, num_dropped = collect_episodes( self.local_evaluator, - self.remote_evaluators, + selected_evaluators or self.remote_evaluators, timeout_seconds=timeout_seconds) orig_episodes = list(episodes) missing = min_history - len(episodes)