mirror of
https://github.com/wassname/ray.git
synced 2026-08-11 05:51:40 +08:00
[rllib] TF model custom_loss() should actually allow access to full rollout data (#4220)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user