From b8436f0f00235f87cc19a1a079018c16d5df683e Mon Sep 17 00:00:00 2001 From: maxco2 Date: Sat, 12 Sep 2020 03:04:44 +0800 Subject: [PATCH] [rllib] Fix SAC and DDPG tensorflow policy can't do `grad_clip` (#10499) * Fix sac_tf_policy clip_by_norm missing argument * Fix ddpg_tf_policy clip_by_norm missing argument * Fix format --- rllib/agents/ddpg/ddpg_tf_policy.py | 4 +++- rllib/agents/sac/sac_tf_policy.py | 4 +++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/rllib/agents/ddpg/ddpg_tf_policy.py b/rllib/agents/ddpg/ddpg_tf_policy.py index 12568e230..3a04bed51 100644 --- a/rllib/agents/ddpg/ddpg_tf_policy.py +++ b/rllib/agents/ddpg/ddpg_tf_policy.py @@ -1,4 +1,5 @@ from gym.spaces import Box +from functools import partial import logging import numpy as np @@ -290,7 +291,8 @@ def gradients_fn(policy, optimizer, loss): # Clip if necessary. if policy.config["grad_clip"]: - clip_func = tf.clip_by_norm + clip_func = partial( + tf.clip_by_norm, clip_norm=policy.config["grad_clip"]) else: clip_func = tf.identity diff --git a/rllib/agents/sac/sac_tf_policy.py b/rllib/agents/sac/sac_tf_policy.py index 26ff292d7..451567547 100644 --- a/rllib/agents/sac/sac_tf_policy.py +++ b/rllib/agents/sac/sac_tf_policy.py @@ -1,4 +1,5 @@ from gym.spaces import Box, Discrete +from functools import partial import logging import ray @@ -320,7 +321,8 @@ def gradients_fn(policy, optimizer, loss): # Clip if necessary. if policy.config["grad_clip"]: - clip_func = tf.clip_by_norm + clip_func = partial( + tf.clip_by_norm, clip_norm=policy.config["grad_clip"]) else: clip_func = tf.identity