mirror of
https://github.com/wassname/ray.git
synced 2026-08-11 05:51:40 +08:00
* Rollback. * Fix issue 7635. * Fix issue 7635. * LINT and bug fix.
This commit is contained in:
@@ -70,14 +70,13 @@ class DistributionalQModel(TFModelV2):
|
||||
action_out = tf.keras.layers.Dense(
|
||||
units=q_hiddens[i],
|
||||
activation_fn=tf.nn.relu,
|
||||
normalizer_fn=tf.keras.layers.LayerNormalization
|
||||
)(action_out)
|
||||
normalizer_fn=tf.keras.layers.LayerNormalization)(
|
||||
action_out)
|
||||
else:
|
||||
action_out = tf.keras.layers.Dense(
|
||||
units=q_hiddens[i],
|
||||
activation=tf.nn.relu,
|
||||
name="hidden_%d" % i
|
||||
)(action_out)
|
||||
name="hidden_%d" % i)(action_out)
|
||||
else:
|
||||
# Avoid postprocessing the outputs. This enables custom models
|
||||
# to be used for parametric action DQN.
|
||||
@@ -92,27 +91,28 @@ class DistributionalQModel(TFModelV2):
|
||||
elif q_hiddens:
|
||||
action_scores = tf.keras.layers.Dense(
|
||||
units=self.action_space.n * num_atoms,
|
||||
activation=None
|
||||
)(action_out)
|
||||
activation=None)(action_out)
|
||||
else:
|
||||
action_scores = model_out
|
||||
|
||||
if num_atoms > 1:
|
||||
# Distributional Q-learning uses a discrete support z
|
||||
# to represent the action value distribution
|
||||
z = tf.range(num_atoms, dtype=tf.float32)
|
||||
z = v_min + z * (v_max - v_min) / float(num_atoms - 1)
|
||||
support_logits_per_action = tf.reshape(
|
||||
tensor=action_scores,
|
||||
shape=(-1, self.action_space.n, num_atoms))
|
||||
support_prob_per_action = tf.nn.softmax(
|
||||
logits=support_logits_per_action)
|
||||
action_scores = tf.reduce_sum(
|
||||
input_tensor=z * support_prob_per_action, axis=-1)
|
||||
logits = support_logits_per_action
|
||||
dist = support_prob_per_action
|
||||
return [
|
||||
action_scores, z, support_logits_per_action, logits, dist
|
||||
]
|
||||
|
||||
def _layer(x):
|
||||
support_logits_per_action = tf.reshape(
|
||||
tensor=x, shape=(-1, self.action_space.n, num_atoms))
|
||||
support_prob_per_action = tf.nn.softmax(
|
||||
logits=support_logits_per_action)
|
||||
x = tf.reduce_sum(
|
||||
input_tensor=z * support_prob_per_action, axis=-1)
|
||||
logits = support_logits_per_action
|
||||
dist = support_prob_per_action
|
||||
return [x, z, support_logits_per_action, logits, dist]
|
||||
|
||||
return tf.keras.layers.Lambda(_layer)(action_scores)
|
||||
else:
|
||||
logits = tf.expand_dims(tf.ones_like(action_scores), -1)
|
||||
dist = tf.expand_dims(tf.ones_like(action_scores), -1)
|
||||
@@ -129,12 +129,10 @@ class DistributionalQModel(TFModelV2):
|
||||
state_out = tf.keras.layers.Dense(
|
||||
units=q_hiddens[i],
|
||||
activation_fn=tf.nn.relu,
|
||||
normalizer_fn=tf.contrib.layers.layer_norm
|
||||
)(state_out)
|
||||
normalizer_fn=tf.contrib.layers.layer_norm)(state_out)
|
||||
else:
|
||||
state_out = tf.keras.layers.Dense(
|
||||
units=q_hiddens[i], activation=tf.nn.relu
|
||||
)(state_out)
|
||||
units=q_hiddens[i], activation=tf.nn.relu)(state_out)
|
||||
if use_noisy:
|
||||
state_score = self._noisy_layer(
|
||||
"dueling_output",
|
||||
@@ -144,8 +142,7 @@ class DistributionalQModel(TFModelV2):
|
||||
non_linear=False)
|
||||
else:
|
||||
state_score = tf.keras.layers.Dense(
|
||||
units=num_atoms, activation=None
|
||||
)(state_out)
|
||||
units=num_atoms, activation=None)(state_out)
|
||||
return state_score
|
||||
|
||||
if tf.executing_eagerly():
|
||||
@@ -187,8 +184,7 @@ class DistributionalQModel(TFModelV2):
|
||||
self.register_variables(self.q_value_head.variables)
|
||||
|
||||
if dueling:
|
||||
state_out = build_state_score_in_scope(
|
||||
self.model_out)
|
||||
state_out = build_state_score_in_scope(self.model_out)
|
||||
self.state_value_head = tf.keras.Model(self.model_out, state_out)
|
||||
self.register_variables(self.state_value_head.variables)
|
||||
|
||||
@@ -256,15 +252,17 @@ class DistributionalQModel(TFModelV2):
|
||||
name=prefix + "_fc_w",
|
||||
shape=[in_size, out_size],
|
||||
dtype=tf.float32,
|
||||
initializer=tf.initializers.GlorotUniform())
|
||||
initializer=tf.initializers.glorot_uniform())
|
||||
b = tf.get_variable(
|
||||
name=prefix + "_fc_b",
|
||||
shape=[out_size],
|
||||
dtype=tf.float32,
|
||||
initializer=tf.zeros_initializer())
|
||||
|
||||
action_activation = tf.nn.xw_plus_b(action_in, w + sigma_w * epsilon_w,
|
||||
b + sigma_b * epsilon_b)
|
||||
action_activation = \
|
||||
tf.keras.layers.Lambda(lambda x: tf.matmul(
|
||||
x, w + sigma_w * epsilon_w) + b + sigma_b * epsilon_b)(
|
||||
action_in)
|
||||
|
||||
if not non_linear:
|
||||
return action_activation
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import numpy as np
|
||||
from tensorflow.python.eager.context import eager_mode
|
||||
import unittest
|
||||
|
||||
import ray.rllib.agents.dqn as dqn
|
||||
@@ -14,21 +15,41 @@ class TestDQN(unittest.TestCase):
|
||||
config = dqn.DEFAULT_CONFIG.copy()
|
||||
config["num_workers"] = 0 # Run locally.
|
||||
|
||||
# tf.
|
||||
config["eager"] = False
|
||||
trainer = dqn.DQNTrainer(config=config, env="CartPole-v0")
|
||||
# Rainbow.
|
||||
rainbow_config = config.copy()
|
||||
rainbow_config["eager"] = False
|
||||
rainbow_config["num_atoms"] = 10
|
||||
rainbow_config["noisy"] = True
|
||||
rainbow_config["double_q"] = True
|
||||
rainbow_config["dueling"] = True
|
||||
rainbow_config["n_step"] = 5
|
||||
trainer = dqn.DQNTrainer(config=rainbow_config, env="CartPole-v0")
|
||||
num_iterations = 2
|
||||
for i in range(num_iterations):
|
||||
results = trainer.train()
|
||||
print(results)
|
||||
|
||||
config["eager"] = True
|
||||
trainer = dqn.DQNTrainer(config=config, env="CartPole-v0")
|
||||
# tf.
|
||||
tf_config = config.copy()
|
||||
tf_config["eager"] = False
|
||||
trainer = dqn.DQNTrainer(config=tf_config, env="CartPole-v0")
|
||||
num_iterations = 2
|
||||
for i in range(num_iterations):
|
||||
results = trainer.train()
|
||||
print(results)
|
||||
|
||||
# Eager.
|
||||
eager_config = config.copy()
|
||||
eager_config["eager"] = True
|
||||
eager_ctx = eager_mode()
|
||||
eager_ctx.__enter__()
|
||||
trainer = dqn.DQNTrainer(config=eager_config, env="CartPole-v0")
|
||||
num_iterations = 2
|
||||
for i in range(num_iterations):
|
||||
results = trainer.train()
|
||||
print(results)
|
||||
eager_ctx.__exit__(None, None, None)
|
||||
|
||||
def test_dqn_exploration_and_soft_q_config(self):
|
||||
"""Tests, whether a DQN Agent outputs exploration/softmaxed actions."""
|
||||
config = dqn.DEFAULT_CONFIG.copy()
|
||||
|
||||
Reference in New Issue
Block a user