mirror of
https://github.com/wassname/ray.git
synced 2026-07-27 11:26:41 +08:00
[rllib] Remove "Common", cleanup some code (#2348)
This commit is contained in:
@@ -9,7 +9,7 @@ import os
|
||||
import pickle
|
||||
|
||||
import tensorflow as tf
|
||||
from ray.rllib.evaluation.common_policy_evaluator import CommonPolicyEvaluator
|
||||
from ray.rllib.evaluation.policy_evaluator import PolicyEvaluator
|
||||
from ray.tune.registry import ENV_CREATOR, _global_registry
|
||||
from ray.tune.result import TrainingResult
|
||||
from ray.tune.trainable import Trainable
|
||||
@@ -115,13 +115,13 @@ class Agent(Trainable):
|
||||
"""Convenience method to return configured local evaluator."""
|
||||
|
||||
return self._make_evaluator(
|
||||
CommonPolicyEvaluator, env_creator, policy_graph, 0)
|
||||
PolicyEvaluator, env_creator, policy_graph, 0)
|
||||
|
||||
def make_remote_evaluators(
|
||||
self, env_creator, policy_graph, count, remote_args):
|
||||
"""Convenience method to return a number of remote evaluators."""
|
||||
|
||||
cls = CommonPolicyEvaluator.as_remote(**remote_args).remote
|
||||
cls = PolicyEvaluator.as_remote(**remote_args).remote
|
||||
return [
|
||||
self._make_evaluator(cls, env_creator, policy_graph, i+1)
|
||||
for i in range(count)]
|
||||
|
||||
@@ -8,11 +8,11 @@ from six.moves import queue
|
||||
import ray
|
||||
from ray.rllib.agents.bc.experience_dataset import ExperienceDataset
|
||||
from ray.rllib.agents.bc.policy import BCPolicy
|
||||
from ray.rllib.evaluation.interface import PolicyEvaluator
|
||||
from ray.rllib.evaluation.interface import EvaluatorInterface
|
||||
from ray.rllib.models import ModelCatalog
|
||||
|
||||
|
||||
class BCEvaluator(PolicyEvaluator):
|
||||
class BCEvaluator(EvaluatorInterface):
|
||||
def __init__(self, env_creator, config, logdir):
|
||||
env = ModelCatalog.get_preprocessor_as_wrapper(env_creator(
|
||||
config["env_config"]), config["model"])
|
||||
|
||||
Reference in New Issue
Block a user