Add better logging for SAC and PPO

This commit is contained in:
Antonin Raffin 2020-03-13 11:43:12 +01:00
parent c39421fa64
commit 29d7018265
5 changed files with 49 additions and 17 deletions

View file

@ -14,6 +14,7 @@ Breaking Changes:
New Features:
^^^^^^^^^^^^^
- Better logging for ``SAC`` and ``PPO``
Bug Fixes:
^^^^^^^^^^

View file

@ -99,6 +99,8 @@ class BaseRLModel(ABC):
# Buffers for logging
self.ep_info_buffer = None # type: Optional[deque]
self.ep_success_buffer = None # type: Optional[deque]
# For logging
self._n_updates = 0 # type: int
# Create and wrap the env if needed
if env is not None:

View file

@ -198,7 +198,7 @@ class PPO(BaseRLModel):
return obs, continue_training
def train(self, gradient_steps: int, batch_size: int = 64) -> None:
def train(self, n_epochs: int, batch_size: int = 64) -> None:
# Update optimizer learning rate
self._update_learning_rate(self.policy.optimizer)
# Compute current clip range
@ -207,9 +207,14 @@ class PPO(BaseRLModel):
if self.clip_range_vf is not None:
clip_range_vf = self.clip_range_vf(self._current_progress)
for gradient_step in range(gradient_steps):
entropy_losses, all_kl_divs = [], []
pg_losses, value_losses = [], []
clip_fractions = []
# train for gradient_steps epochs
for epoch in range(n_epochs):
approx_kl_divs = []
# Sample replay buffer
# Do a complete pass on the rollout buffer
for rollout_data in self.rollout_buffer.get(batch_size):
actions = rollout_data.actions
@ -236,6 +241,11 @@ class PPO(BaseRLModel):
policy_loss_2 = advantages * th.clamp(ratio, 1 - clip_range, 1 + clip_range)
policy_loss = -th.min(policy_loss_1, policy_loss_2).mean()
# Logging
pg_losses.append(policy_loss.item())
clip_fraction = th.mean((th.abs(ratio - 1) > clip_range).float()).item()
clip_fractions.append(clip_fraction)
if self.clip_range_vf is None:
# No clipping
values_pred = values
@ -246,6 +256,7 @@ class PPO(BaseRLModel):
clip_range_vf)
# Value loss using the TD(gae_lambda) target
value_loss = F.mse_loss(rollout_data.returns, values_pred)
value_losses.append(value_loss.item())
# Entropy loss favor exploration
if entropy is None:
@ -254,6 +265,8 @@ class PPO(BaseRLModel):
else:
entropy_loss = -th.mean(entropy)
entropy_losses.append(entropy_loss.item())
loss = policy_loss + self.ent_coef * entropy_loss + self.vf_coef * value_loss
# Optimization step
@ -264,23 +277,27 @@ class PPO(BaseRLModel):
self.policy.optimizer.step()
approx_kl_divs.append(th.mean(rollout_data.old_log_prob - log_prob).detach().cpu().numpy())
all_kl_divs.append(np.mean(approx_kl_divs))
if self.target_kl is not None and np.mean(approx_kl_divs) > 1.5 * self.target_kl:
print("Early stopping at step {} due to reaching max kl: {:.2f}".format(gradient_step,
np.mean(approx_kl_divs)))
print(f"Early stopping at step {epoch} due to reaching max kl: {np.mean(approx_kl_divs):.2f}")
break
self._n_updates += n_epochs
explained_var = explained_variance(self.rollout_buffer.returns.flatten(),
self.rollout_buffer.values.flatten())
logger.logkv("n_updates", self._n_updates)
logger.logkv("clip_fraction", np.mean(clip_fraction))
logger.logkv("clip_range", clip_range)
if self.clip_range_vf is not None:
logger.logkv("clip_range_vf", clip_range_vf)
logger.logkv("approx_kl", np.mean(approx_kl_divs))
logger.logkv("explained_variance", explained_var)
# TODO: gather stats for the entropy and other losses?
logger.logkv("entropy_loss", entropy_loss.item())
logger.logkv("policy_loss", policy_loss.item())
logger.logkv("value_loss", value_loss.item())
logger.logkv("entropy_loss", np.mean(entropy_losses))
logger.logkv("policy_gradient_loss", np.mean(pg_losses))
logger.logkv("value_loss", np.mean(value_losses))
if hasattr(self.policy, 'log_std'):
logger.logkv("std", th.exp(self.policy.log_std).mean().item())

View file

@ -173,8 +173,8 @@ class SAC(OffPolicyRLModel):
self._update_learning_rate(optimizers)
ent_coef_loss, ent_coef = th.zeros(1), th.zeros(1)
actor_loss, critic_loss = th.zeros(1), th.zeros(1)
ent_coef_losses, ent_coefs = [], []
actor_losses, critic_losses = [], []
for gradient_step in range(gradient_steps):
# Sample replay buffer
@ -195,9 +195,12 @@ class SAC(OffPolicyRLModel):
# see https://github.com/rail-berkeley/softlearning/issues/60
ent_coef = th.exp(self.log_ent_coef.detach())
ent_coef_loss = -(self.log_ent_coef * (log_prob + self.target_entropy).detach()).mean()
ent_coef_losses.append(ent_coef_loss.item())
else:
ent_coef = self.ent_coef_tensor
ent_coefs.append(ent_coef.item())
# Optimize entropy coefficient, also called
# entropy temperature or alpha in the paper
if ent_coef_loss is not None:
@ -221,6 +224,7 @@ class SAC(OffPolicyRLModel):
# Compute critic loss
critic_loss = 0.5 * (F.mse_loss(current_q1, q_backup) + F.mse_loss(current_q2, q_backup))
critic_losses.append(critic_loss.item())
# Optimize the critic
self.critic.optimizer.zero_grad()
@ -232,6 +236,7 @@ class SAC(OffPolicyRLModel):
qf1_pi, qf2_pi = self.critic.forward(replay_data.observations, actions_pi)
min_qf_pi = th.min(qf1_pi, qf2_pi)
actor_loss = (ent_coef * log_prob - min_qf_pi).mean()
actor_losses.append(actor_loss.item())
# Optimize the actor
self.actor.optimizer.zero_grad()
@ -243,12 +248,14 @@ class SAC(OffPolicyRLModel):
for param, target_param in zip(self.critic.parameters(), self.critic_target.parameters()):
target_param.data.copy_(self.tau * param.data + (1 - self.tau) * target_param.data)
# TODO: average
logger.logkv("ent_coef", ent_coef.item())
logger.logkv("actor_loss", actor_loss.item())
logger.logkv("critic_loss", critic_loss.item())
if ent_coef_loss is not None:
logger.logkv("ent_coef_loss", ent_coef_loss.item())
self._n_updates += gradient_steps
logger.logkv("n_updates", self._n_updates)
logger.logkv("ent_coef", np.mean(ent_coefs))
logger.logkv("actor_loss", np.mean(actor_losses))
logger.logkv("critic_loss", np.mean(critic_losses))
if len(ent_coef_losses) > 0:
logger.logkv("ent_coef_loss", np.mean(ent_coef_losses))
def learn(self,
total_timesteps: int,

View file

@ -2,6 +2,7 @@ import torch as th
import torch.nn.functional as F
from typing import List, Tuple, Type, Union, Callable, Optional, Dict, Any
from torchy_baselines.common import logger
from torchy_baselines.common.base_class import OffPolicyRLModel
from torchy_baselines.common.buffers import ReplayBuffer
from torchy_baselines.common.noise import ActionNoise
@ -215,6 +216,10 @@ class TD3(OffPolicyRLModel):
if gradient_step % policy_delay == 0:
self.train_actor(replay_data=replay_data, tau_actor=self.tau, tau_critic=self.tau)
self._n_updates += gradient_steps
logger.logkv("n_updates", self._n_updates)
def train_sde(self) -> None:
# Update optimizer learning rate
# self._update_learning_rate(self.policy.optimizer)