[rllib] TF model custom_loss() should actually allow access to full rollout data (#4220)

This commit is contained in:
Eric Liang
2019-03-02 22:57:51 -08:00
committed by GitHub
parent ff6dd8459a
commit ba03048254
6 changed files with 27 additions and 18 deletions
@@ -333,13 +333,6 @@ class DDPGPolicyGraph(TFPolicyGraph):
self.loss.critic_loss += (
config["l2_reg"] * 0.5 * tf.nn.l2_loss(var))
# Model self-supervised losses
self.loss.actor_loss = self.p_model.custom_loss(self.loss.actor_loss)
self.loss.critic_loss = self.q_model.custom_loss(self.loss.critic_loss)
if self.config["twin_q"]:
self.loss.critic_loss = self.twin_q_model.custom_loss(
self.loss.critic_loss)
# update_target_fn will be called periodically to copy Q network to
# target Q network
self.tau_value = config.get("tau")
@@ -375,6 +368,17 @@ class DDPGPolicyGraph(TFPolicyGraph):
("dones", self.done_mask),
("weights", self.importance_weights),
]
input_dict = dict(self.loss_inputs)
# Model self-supervised losses
self.loss.actor_loss = self.p_model.custom_loss(
self.loss.actor_loss, input_dict)
self.loss.critic_loss = self.q_model.custom_loss(
self.loss.critic_loss, input_dict)
if self.config["twin_q"]:
self.loss.critic_loss = self.twin_q_model.custom_loss(
self.loss.critic_loss, input_dict)
TFPolicyGraph.__init__(
self,
observation_space,