mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-30 20:18:15 +00:00
Enable logger for SAC/TD3 + refactor
This commit is contained in:
parent
dbaa5daca6
commit
b5656531d1
7 changed files with 100 additions and 60 deletions
|
|
@ -55,15 +55,10 @@ class CEMRL(TD3):
|
|||
pop_size=self.pop_size, antithetic=not self.pop_size % 2, parents=self.pop_size // 2,
|
||||
elitism=self.elitism)
|
||||
|
||||
def learn(self, total_timesteps, callback=None, log_interval=100,
|
||||
def learn(self, total_timesteps, callback=None, log_interval=4,
|
||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="CEMRL", reset_num_timesteps=True):
|
||||
|
||||
timesteps_since_eval, actor_steps = 0, 0
|
||||
episode_num = 0
|
||||
evaluations = []
|
||||
start_time = time.time()
|
||||
eval_env = self._get_eval_env(eval_env)
|
||||
obs = self.env.reset()
|
||||
timesteps_since_eval, episode_num, evaluations, obs, eval_env = self._setup_learn(eval_env)
|
||||
|
||||
while self.num_timesteps < total_timesteps:
|
||||
|
||||
|
|
@ -127,7 +122,7 @@ class CEMRL(TD3):
|
|||
|
||||
if self.verbose > 0:
|
||||
print("Eval num_timesteps={}, mean_reward={:.2f}".format(self.num_timesteps, evaluations[-1]))
|
||||
print("FPS: {:.2f}".format(self.num_timesteps / (time.time() - start_time)))
|
||||
print("FPS: {:.2f}".format(self.num_timesteps / (time.time() - self.start_time)))
|
||||
|
||||
actor_steps = 0
|
||||
# evaluate all actors
|
||||
|
|
@ -141,7 +136,8 @@ class CEMRL(TD3):
|
|||
learning_starts=self.learning_starts,
|
||||
num_timesteps=self.num_timesteps,
|
||||
replay_buffer=self.replay_buffer,
|
||||
obs=obs)
|
||||
obs=obs, episode_num=episode_num,
|
||||
log_interval=log_interval)
|
||||
|
||||
# Unpack
|
||||
episode_reward, episode_timesteps, n_episodes, obs = rollout
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
import time
|
||||
from abc import ABCMeta, abstractmethod
|
||||
from collections import deque
|
||||
|
||||
import gym
|
||||
import torch as th
|
||||
|
|
@ -7,6 +9,8 @@ import numpy as np
|
|||
from torchy_baselines.common.policies import get_policy_from_name
|
||||
from torchy_baselines.common.utils import set_random_seed
|
||||
from torchy_baselines.common.vec_env import DummyVecEnv, VecEnv
|
||||
from torchy_baselines.common.monitor import Monitor
|
||||
from torchy_baselines.common import logger
|
||||
|
||||
|
||||
class BaseRLModel(object):
|
||||
|
|
@ -21,11 +25,14 @@ class BaseRLModel(object):
|
|||
:param device: (str or th.device) Device on which the code should.
|
||||
By default, it will try to use a Cuda compatible device and fallback to cpu
|
||||
if it is not possible.
|
||||
:param monitor_wrapper: (bool) When creating an environment, whether to wrap it
|
||||
or not in a Monitor wrapper.
|
||||
"""
|
||||
__metaclass__ = ABCMeta
|
||||
|
||||
def __init__(self, policy, env, policy_base, policy_kwargs=None,
|
||||
verbose=0, device='auto', support_multi_env=False, create_eval_env=False):
|
||||
verbose=0, device='auto', support_multi_env=False,
|
||||
create_eval_env=False, monitor_wrapper=True, seed=None):
|
||||
if isinstance(policy, str) and policy_base is not None:
|
||||
self.policy = get_policy_from_name(policy_base, policy)
|
||||
else:
|
||||
|
|
@ -48,14 +55,22 @@ class BaseRLModel(object):
|
|||
self.params = None
|
||||
self.eval_env = None
|
||||
self.replay_buffer = None
|
||||
self.seed = seed
|
||||
|
||||
if env is not None:
|
||||
if isinstance(env, str):
|
||||
if create_eval_env:
|
||||
self.eval_env = DummyVecEnv([lambda: gym.make(env)])
|
||||
eval_env = gym.make(env)
|
||||
if monitor_wrapper:
|
||||
eval_env = Monitor(eval_env, filename=None)
|
||||
self.eval_env = DummyVecEnv([lambda: eval_env])
|
||||
if self.verbose >= 1:
|
||||
print("Creating environment from the given name, wrapped in a DummyVecEnv.")
|
||||
env = DummyVecEnv([lambda: gym.make(env)])
|
||||
|
||||
env = gym.make(env)
|
||||
if monitor_wrapper:
|
||||
env = Monitor(env, filename=None)
|
||||
env = DummyVecEnv([lambda: env])
|
||||
|
||||
self.observation_space = env.observation_space
|
||||
self.action_space = env.action_space
|
||||
|
|
@ -96,6 +111,17 @@ class BaseRLModel(object):
|
|||
low, high = self.action_space.low, self.action_space.high
|
||||
return low + (0.5 * (scaled_action + 1.0) * (high - low))
|
||||
|
||||
@staticmethod
|
||||
def safe_mean(arr):
|
||||
"""
|
||||
Compute the mean of an array if there is at least one element.
|
||||
For empty array, return nan. It is used for logging only.
|
||||
|
||||
:param arr: (np.ndarray)
|
||||
:return: (float)
|
||||
"""
|
||||
return np.nan if len(arr) == 0 else np.mean(arr)
|
||||
|
||||
def get_env(self):
|
||||
"""
|
||||
returns the current environment (can be None if not defined)
|
||||
|
|
@ -231,10 +257,24 @@ class BaseRLModel(object):
|
|||
if self.eval_env is not None:
|
||||
self.eval_env.seed(seed)
|
||||
|
||||
def _setup_learn(self, eval_env):
|
||||
self.start_time = time.time()
|
||||
self.ep_info_buffer = deque(maxlen=100)
|
||||
if self.action_noise is not None:
|
||||
self.action_noise.reset()
|
||||
timesteps_since_eval, episode_num = 0, 0
|
||||
evaluations = []
|
||||
if eval_env is not None and self.seed is not None:
|
||||
eval_env.seed(self.seed)
|
||||
eval_env = self._get_eval_env(eval_env)
|
||||
obs = self.env.reset()
|
||||
return timesteps_since_eval, episode_num, evaluations, obs, eval_env
|
||||
|
||||
def collect_rollouts(self, env, n_episodes=1, n_steps=-1, action_noise=None,
|
||||
deterministic=False, callback=None,
|
||||
learning_starts=0, num_timesteps=0,
|
||||
replay_buffer=None, obs=None):
|
||||
replay_buffer=None, obs=None,
|
||||
episode_num=0, log_interval=None):
|
||||
|
||||
episode_rewards = []
|
||||
total_timesteps = []
|
||||
|
|
@ -264,11 +304,17 @@ class BaseRLModel(object):
|
|||
action = np.clip(action + action_noise(), -1, 1)
|
||||
|
||||
# Rescale and perform action
|
||||
new_obs, reward, done, _ = env.step(self.unscale_action(action))
|
||||
new_obs, reward, done, infos = env.step(self.unscale_action(action))
|
||||
|
||||
done_bool = [float(done[0])]
|
||||
episode_reward += reward
|
||||
|
||||
# Retrieve reward and episode length if using Monitor wrapper
|
||||
for info in infos:
|
||||
maybe_ep_info = info.get('episode')
|
||||
if maybe_ep_info is not None:
|
||||
self.ep_info_buffer.extend([maybe_ep_info])
|
||||
|
||||
# Store data in replay buffer
|
||||
if replay_buffer is not None:
|
||||
replay_buffer.add(obs, new_obs, action, reward, done_bool)
|
||||
|
|
@ -288,6 +334,21 @@ class BaseRLModel(object):
|
|||
if action_noise is not None:
|
||||
action_noise.reset()
|
||||
|
||||
# Display training infos
|
||||
if self.verbose >= 1 and log_interval is not None and (episode_num + total_episodes) % log_interval == 0:
|
||||
fps = int(num_timesteps / (time.time() - self.start_time))
|
||||
logger.logkv("episodes", episode_num + total_episodes)
|
||||
# logger.logkv("mean 100 episode reward", mean_reward)
|
||||
if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0:
|
||||
logger.logkv('ep_rew_mean', self.safe_mean([ep_info['r'] for ep_info in self.ep_info_buffer]))
|
||||
logger.logkv('ep_len_mean', self.safe_mean([ep_info['l'] for ep_info in self.ep_info_buffer]))
|
||||
# logger.logkv("n_updates", n_updates)
|
||||
# logger.logkv("current_lr", current_lr)
|
||||
logger.logkv("fps", fps)
|
||||
logger.logkv('time_elapsed', int(time.time() - self.start_time))
|
||||
logger.logkv("total timesteps", num_timesteps)
|
||||
logger.dumpkvs()
|
||||
|
||||
mean_reward = np.mean(episode_rewards) if total_episodes > 0 else 0.0
|
||||
|
||||
return mean_reward, total_steps, total_episodes, obs
|
||||
|
|
|
|||
|
|
@ -74,10 +74,9 @@ class PPO(BaseRLModel):
|
|||
|
||||
super(PPO, self).__init__(policy, env, PPOPolicy, policy_kwargs=policy_kwargs,
|
||||
verbose=verbose, device=device,
|
||||
create_eval_env=create_eval_env, support_multi_env=True)
|
||||
create_eval_env=create_eval_env, support_multi_env=True, seed=seed)
|
||||
|
||||
self.learning_rate = learning_rate
|
||||
self.seed = seed
|
||||
self.batch_size = batch_size
|
||||
self.n_epochs = n_epochs
|
||||
self.n_steps = n_steps
|
||||
|
|
|
|||
|
|
@ -10,12 +10,9 @@ LOG_STD_MIN = -20
|
|||
|
||||
|
||||
class Actor(BaseNetwork):
|
||||
def __init__(self, obs_dim, action_dim, net_arch=None, activation_fn=nn.ReLU):
|
||||
def __init__(self, obs_dim, action_dim, net_arch, activation_fn=nn.ReLU):
|
||||
super(Actor, self).__init__()
|
||||
|
||||
if net_arch is None:
|
||||
net_arch = [256, 256]
|
||||
|
||||
# TODO: orthogonal initialization?
|
||||
actor_net = create_mlp(obs_dim, -1, net_arch, activation_fn)
|
||||
self.actor_net = nn.Sequential(*actor_net)
|
||||
|
|
@ -45,12 +42,9 @@ class Actor(BaseNetwork):
|
|||
|
||||
class Critic(BaseNetwork):
|
||||
def __init__(self, obs_dim, action_dim,
|
||||
net_arch=None, activation_fn=nn.ReLU):
|
||||
net_arch, activation_fn=nn.ReLU):
|
||||
super(Critic, self).__init__()
|
||||
|
||||
if net_arch is None:
|
||||
net_arch = [256, 256]
|
||||
|
||||
q1_net = create_mlp(obs_dim + action_dim, 1, net_arch, activation_fn)
|
||||
self.q1_net = nn.Sequential(*q1_net)
|
||||
|
||||
|
|
@ -72,6 +66,10 @@ class SACPolicy(BasePolicy):
|
|||
learning_rate=3e-4, net_arch=None, device='cpu',
|
||||
activation_fn=nn.ReLU):
|
||||
super(SACPolicy, self).__init__(observation_space, action_space, device)
|
||||
|
||||
if net_arch is None:
|
||||
net_arch = [256, 256]
|
||||
|
||||
self.obs_dim = self.observation_space.shape[0]
|
||||
self.action_dim = self.action_space.shape[0]
|
||||
self.net_arch = net_arch
|
||||
|
|
|
|||
|
|
@ -62,10 +62,9 @@ class SAC(BaseRLModel):
|
|||
_init_setup_model=True):
|
||||
|
||||
super(SAC, self).__init__(policy, env, SACPolicy, policy_kwargs, verbose, device,
|
||||
create_eval_env=create_eval_env)
|
||||
create_eval_env=create_eval_env, seed=seed)
|
||||
|
||||
self.learning_rate = learning_rate
|
||||
self.seed = seed
|
||||
self.target_entropy = target_entropy
|
||||
self.log_ent_coef = None
|
||||
self.target_update_interval = target_update_interval
|
||||
|
|
@ -90,7 +89,8 @@ class SAC(BaseRLModel):
|
|||
|
||||
def _setup_model(self):
|
||||
obs_dim, action_dim = self.observation_space.shape[0], self.action_space.shape[0]
|
||||
self.set_random_seed(self.seed)
|
||||
if self.seed is not None:
|
||||
self.set_random_seed(self.seed)
|
||||
|
||||
# Target entropy is used when learning the entropy coefficient
|
||||
if self.target_entropy == 'auto':
|
||||
|
|
@ -218,15 +218,11 @@ class SAC(BaseRLModel):
|
|||
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)
|
||||
|
||||
def learn(self, total_timesteps, callback=None, log_interval=100,
|
||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="SAC", reset_num_timesteps=True):
|
||||
def learn(self, total_timesteps, callback=None, log_interval=4,
|
||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="SAC",
|
||||
reset_num_timesteps=True):
|
||||
|
||||
timesteps_since_eval = 0
|
||||
episode_num = 0
|
||||
evaluations = []
|
||||
start_time = time.time()
|
||||
eval_env = self._get_eval_env(eval_env)
|
||||
obs = self.env.reset()
|
||||
timesteps_since_eval, episode_num, evaluations, obs, eval_env = self._setup_learn(eval_env)
|
||||
|
||||
while self.num_timesteps < total_timesteps:
|
||||
|
||||
|
|
@ -241,7 +237,8 @@ class SAC(BaseRLModel):
|
|||
learning_starts=self.learning_starts,
|
||||
num_timesteps=self.num_timesteps,
|
||||
replay_buffer=self.replay_buffer,
|
||||
obs=obs)
|
||||
obs=obs, episode_num=episode_num,
|
||||
log_interval=log_interval)
|
||||
# Unpack
|
||||
episode_reward, episode_timesteps, n_episodes, obs = rollout
|
||||
|
||||
|
|
@ -264,7 +261,7 @@ class SAC(BaseRLModel):
|
|||
evaluations.append(mean_reward)
|
||||
if self.verbose > 0:
|
||||
print("Eval num_timesteps={}, mean_reward={:.2f}".format(self.num_timesteps, evaluations[-1]))
|
||||
print("FPS: {:.2f}".format(self.num_timesteps / (time.time() - start_time)))
|
||||
print("FPS: {:.2f}".format(self.num_timesteps / (time.time() - self.start_time)))
|
||||
|
||||
return self
|
||||
|
||||
|
|
|
|||
|
|
@ -5,12 +5,9 @@ from torchy_baselines.common.policies import BasePolicy, register_policy, create
|
|||
|
||||
|
||||
class Actor(BaseNetwork):
|
||||
def __init__(self, obs_dim, action_dim, net_arch=None, activation_fn=nn.ReLU):
|
||||
def __init__(self, obs_dim, action_dim, net_arch, activation_fn=nn.ReLU):
|
||||
super(Actor, self).__init__()
|
||||
|
||||
if net_arch is None:
|
||||
net_arch = [400, 300]
|
||||
|
||||
# TODO: orthogonal initialization?
|
||||
actor_net = create_mlp(obs_dim, action_dim, net_arch, activation_fn, squash_out=True)
|
||||
self.actor_net = nn.Sequential(*actor_net)
|
||||
|
|
@ -21,12 +18,9 @@ class Actor(BaseNetwork):
|
|||
|
||||
class Critic(BaseNetwork):
|
||||
def __init__(self, obs_dim, action_dim,
|
||||
net_arch=None, activation_fn=nn.ReLU):
|
||||
net_arch, activation_fn=nn.ReLU):
|
||||
super(Critic, self).__init__()
|
||||
|
||||
if net_arch is None:
|
||||
net_arch = [400, 300]
|
||||
|
||||
q1_net = create_mlp(obs_dim + action_dim, 1, net_arch, activation_fn)
|
||||
self.q1_net = nn.Sequential(*q1_net)
|
||||
|
||||
|
|
@ -48,6 +42,10 @@ class TD3Policy(BasePolicy):
|
|||
learning_rate=1e-3, net_arch=None, device='cpu',
|
||||
activation_fn=nn.ReLU):
|
||||
super(TD3Policy, self).__init__(observation_space, action_space, device)
|
||||
|
||||
if net_arch is None:
|
||||
net_arch = [400, 300]
|
||||
|
||||
self.obs_dim = self.observation_space.shape[0]
|
||||
self.action_dim = self.action_space.shape[0]
|
||||
self.net_arch = net_arch
|
||||
|
|
|
|||
|
|
@ -54,10 +54,7 @@ class TD3(BaseRLModel):
|
|||
seed=0, device='auto', _init_setup_model=True):
|
||||
|
||||
super(TD3, self).__init__(policy, env, TD3Policy, policy_kwargs, verbose, device,
|
||||
create_eval_env=create_eval_env)
|
||||
|
||||
self.buffer_size = buffer_size
|
||||
self.seed = seed
|
||||
create_eval_env=create_eval_env, seed=seed)
|
||||
|
||||
self.buffer_size = buffer_size
|
||||
# TODO: accept callables
|
||||
|
|
@ -185,17 +182,10 @@ class TD3(BaseRLModel):
|
|||
if gradient_step % policy_delay == 0:
|
||||
self.train_actor(replay_data=replay_data, tau_actor=self.tau, tau_critic=self.tau)
|
||||
|
||||
def learn(self, total_timesteps, callback=None, log_interval=100,
|
||||
def learn(self, total_timesteps, callback=None, log_interval=4,
|
||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="TD3", reset_num_timesteps=True):
|
||||
|
||||
timesteps_since_eval = 0
|
||||
episode_num = 0
|
||||
evaluations = []
|
||||
start_time = time.time()
|
||||
eval_env = self._get_eval_env(eval_env)
|
||||
obs = self.env.reset()
|
||||
if self.action_noise is not None:
|
||||
self.action_noise.reset()
|
||||
timesteps_since_eval, episode_num, evaluations, obs, eval_env = self._setup_learn(eval_env)
|
||||
|
||||
while self.num_timesteps < total_timesteps:
|
||||
|
||||
|
|
@ -210,7 +200,8 @@ class TD3(BaseRLModel):
|
|||
learning_starts=self.learning_starts,
|
||||
num_timesteps=self.num_timesteps,
|
||||
replay_buffer=self.replay_buffer,
|
||||
obs=obs)
|
||||
obs=obs, episode_num=episode_num,
|
||||
log_interval=log_interval)
|
||||
# Unpack
|
||||
episode_reward, episode_timesteps, n_episodes, obs = rollout
|
||||
|
||||
|
|
@ -233,7 +224,7 @@ class TD3(BaseRLModel):
|
|||
evaluations.append(mean_reward)
|
||||
if self.verbose > 0:
|
||||
print("Eval num_timesteps={}, mean_reward={:.2f}".format(self.num_timesteps, evaluations[-1]))
|
||||
print("FPS: {:.2f}".format(self.num_timesteps / (time.time() - start_time)))
|
||||
print("FPS: {:.2f}".format(self.num_timesteps / (time.time() - self.start_time)))
|
||||
|
||||
return self
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue