[RLlib] Add grad_clip config option to MARWIL and stabilize grad clipping against inf global_norms. (#13634)

This commit is contained in:
Sven Mika
2021-01-22 19:36:02 +01:00
committed by GitHub
parent da5928304a
commit d629292d63
4 changed files with 15 additions and 4 deletions
+2
View File
@@ -21,6 +21,8 @@ DEFAULT_CONFIG = with_common_config({
"beta": 1.0,
# Balancing value estimation loss and policy optimization loss.
"vf_coeff": 1.0,
# If specified, clip the global norm of gradients by this amount.
"grad_clip": None,
# Whether to calculate cumulative rewards.
"postprocess_inputs": True,
# Whether to rollout "complete_episodes" or "truncate_episodes".
+3 -1
View File
@@ -1,6 +1,7 @@
import logging
import ray
from ray.rllib.agents.ppo.ppo_tf_policy import compute_and_clip_gradients
from ray.rllib.policy.sample_batch import SampleBatch
from ray.rllib.evaluation.postprocessing import compute_advantages, \
Postprocessing
@@ -133,7 +134,7 @@ class MARWILLoss:
# Exponentially weighted advantages.
c = tf.math.sqrt(policy._moving_average_sqd_adv_norm)
exp_advs = tf.math.exp(beta * (adv / c))
exp_advs = tf.math.exp(beta * (adv / (1e-8 + c)))
# Static graph.
else:
update_adv_norm = tf1.assign_add(
@@ -200,4 +201,5 @@ MARWILTFPolicy = build_tf_policy(
stats_fn=stats,
postprocess_fn=postprocess_advantages,
before_loss_init=setup_mixins,
gradients_fn=compute_and_clip_gradients,
mixins=[ValueNetworkMixin])
+2 -1
View File
@@ -4,7 +4,7 @@ from ray.rllib.evaluation.postprocessing import Postprocessing
from ray.rllib.policy.policy_template import build_policy_class
from ray.rllib.policy.sample_batch import SampleBatch
from ray.rllib.utils.framework import try_import_torch
from ray.rllib.utils.torch_ops import explained_variance
from ray.rllib.utils.torch_ops import apply_grad_clipping, explained_variance
torch, _ = try_import_torch()
@@ -98,5 +98,6 @@ MARWILTorchPolicy = build_policy_class(
get_default_config=lambda: ray.rllib.agents.marwil.marwil.DEFAULT_CONFIG,
stats_fn=stats,
postprocess_fn=postprocess_advantages,
extra_grad_process_fn=apply_grad_clipping,
before_loss_init=setup_mixins,
mixins=[ValueNetworkMixin])
+8 -2
View File
@@ -182,9 +182,15 @@ def compute_and_clip_gradients(policy: Policy, optimizer: LocalOptimizer,
# Clip by global norm, if necessary.
if policy.config["grad_clip"] is not None:
# Defuse inf gradients (due to super large losses).
grads = [g for (g, v) in grads_and_vars]
policy.grads, _ = tf.clip_by_global_norm(grads,
policy.config["grad_clip"])
grads, _ = tf.clip_by_global_norm(grads, policy.config["grad_clip"])
# If the global_norm is inf -> All grads will be NaN. Stabilize this
# here by setting them to 0.0. This will simply ignore destructive loss
# calculations.
policy.grads = [
tf.where(tf.math.is_nan(g), tf.zeros_like(g), g) for g in grads
]
clipped_grads_and_vars = list(zip(policy.grads, variables))
return clipped_grads_and_vars
else: