mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-30 20:18:15 +00:00
Add better logging for SAC and PPO
This commit is contained in:
parent
c39421fa64
commit
29d7018265
5 changed files with 49 additions and 17 deletions
|
|
@ -14,6 +14,7 @@ Breaking Changes:
|
|||
|
||||
New Features:
|
||||
^^^^^^^^^^^^^
|
||||
- Better logging for ``SAC`` and ``PPO``
|
||||
|
||||
Bug Fixes:
|
||||
^^^^^^^^^^
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in a new issue