mirror of
https://github.com/wassname/ray.git
synced 2026-07-23 13:10:11 +08:00
Move TensorFlowVariables to ray.experimental.tf_utils. (#4145)
This commit is contained in:
committed by
Philipp Moritz
parent
615d5516d1
commit
7b04ed059e
@@ -10,6 +10,7 @@ import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
import ray
|
||||
import ray.experimental.tf_utils
|
||||
from ray.rllib.evaluation.sampler import _unbatch_tuple_actions
|
||||
from ray.rllib.utils.filter import get_filter
|
||||
from ray.rllib.models import ModelCatalog
|
||||
@@ -81,7 +82,7 @@ class GenericPolicy(object):
|
||||
dist = dist_class(model.outputs)
|
||||
self.sampler = dist.sample()
|
||||
|
||||
self.variables = ray.experimental.TensorFlowVariables(
|
||||
self.variables = ray.experimental.tf_utils.TensorFlowVariables(
|
||||
model.outputs, self.sess)
|
||||
|
||||
self.num_params = sum(
|
||||
|
||||
@@ -8,8 +8,9 @@ import tensorflow as tf
|
||||
import tensorflow.contrib.layers as layers
|
||||
|
||||
import ray
|
||||
from ray.rllib.agents.dqn.dqn_policy_graph import _huber_loss, \
|
||||
_minimize_and_clip, _scope_vars, _postprocess_dqn
|
||||
import ray.experimental.tf_utils
|
||||
from ray.rllib.agents.dqn.dqn_policy_graph import (
|
||||
_huber_loss, _minimize_and_clip, _scope_vars, _postprocess_dqn)
|
||||
from ray.rllib.models import ModelCatalog
|
||||
from ray.rllib.utils.annotations import override
|
||||
from ray.rllib.utils.error import UnsupportedSpaceException
|
||||
@@ -387,7 +388,7 @@ class DDPGPolicyGraph(TFPolicyGraph):
|
||||
|
||||
# Note that this encompasses both the policy and Q-value networks and
|
||||
# their corresponding target networks
|
||||
self.variables = ray.experimental.TensorFlowVariables(
|
||||
self.variables = ray.experimental.tf_utils.TensorFlowVariables(
|
||||
tf.group(q_tp0, q_tp1), self.sess)
|
||||
|
||||
# Hard initial update
|
||||
|
||||
@@ -10,6 +10,7 @@ import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
import ray
|
||||
import ray.experimental.tf_utils
|
||||
from ray.rllib.evaluation.sampler import _unbatch_tuple_actions
|
||||
from ray.rllib.models import ModelCatalog
|
||||
from ray.rllib.utils.filter import get_filter
|
||||
@@ -59,7 +60,7 @@ class GenericPolicy(object):
|
||||
dist = dist_class(model.outputs)
|
||||
self.sampler = dist.sample()
|
||||
|
||||
self.variables = ray.experimental.TensorFlowVariables(
|
||||
self.variables = ray.experimental.tf_utils.TensorFlowVariables(
|
||||
model.outputs, self.sess)
|
||||
|
||||
self.num_params = sum(
|
||||
|
||||
Reference in New Issue
Block a user