mirror of
https://github.com/wassname/ray.git
synced 2026-07-20 12:40:20 +08:00
[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
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user