mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-15 22:10:25 +00:00
Refactor evaluation
This commit is contained in:
parent
d514cd9126
commit
a628354721
8 changed files with 125 additions and 143 deletions
|
|
@ -130,8 +130,10 @@ class A2C(PPO):
|
||||||
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
||||||
|
|
||||||
def learn(self, total_timesteps, callback=None, log_interval=100,
|
def learn(self, total_timesteps, callback=None, log_interval=100,
|
||||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="A2C", reset_num_timesteps=True):
|
eval_env=None, eval_freq=-1, n_eval_episodes=5,
|
||||||
|
tb_log_name="A2C", eval_log_path=None, reset_num_timesteps=True):
|
||||||
|
|
||||||
return super(A2C, self).learn(total_timesteps=total_timesteps, callback=callback, log_interval=log_interval,
|
return super(A2C, self).learn(total_timesteps=total_timesteps, callback=callback, log_interval=log_interval,
|
||||||
eval_env=eval_env, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes,
|
eval_env=eval_env, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes,
|
||||||
tb_log_name=tb_log_name, reset_num_timesteps=reset_num_timesteps)
|
tb_log_name=tb_log_name, eval_log_path=eval_log_path,
|
||||||
|
reset_num_timesteps=reset_num_timesteps)
|
||||||
|
|
|
||||||
|
|
@ -100,9 +100,10 @@ class CEMRL(TD3):
|
||||||
elitism=self.elitism)
|
elitism=self.elitism)
|
||||||
|
|
||||||
def learn(self, total_timesteps, callback=None, log_interval=4,
|
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):
|
eval_env=None, eval_freq=-1, n_eval_episodes=5,
|
||||||
|
tb_log_name="CEMRL", eval_log_path=None, reset_num_timesteps=True):
|
||||||
|
|
||||||
timesteps_since_eval, episode_num, evaluations, obs, eval_env, callback = self._setup_learn(eval_env, callback)
|
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq, n_eval_episodes, eval_log_path)
|
||||||
actor_steps = 0
|
actor_steps = 0
|
||||||
continue_training = True
|
continue_training = True
|
||||||
|
|
||||||
|
|
@ -155,21 +156,6 @@ class CEMRL(TD3):
|
||||||
# Get the params back in the population
|
# Get the params back in the population
|
||||||
self.es_params[i] = self.actor.parameters_to_vector()
|
self.es_params[i] = self.actor.parameters_to_vector()
|
||||||
|
|
||||||
# Evaluate agent
|
|
||||||
if 0 < eval_freq <= timesteps_since_eval and eval_env is not None:
|
|
||||||
timesteps_since_eval %= eval_freq
|
|
||||||
|
|
||||||
self.actor.load_from_vector(self.es.mu)
|
|
||||||
sync_envs_normalization(self.env, eval_env)
|
|
||||||
|
|
||||||
mean_reward, std_reward = evaluate_policy(self, eval_env, n_eval_episodes)
|
|
||||||
evaluations.append(mean_reward)
|
|
||||||
|
|
||||||
if self.verbose > 0:
|
|
||||||
print("Eval num_timesteps={}, "
|
|
||||||
"episode_reward={:.2f} +/- {:.2f}".format(self.num_timesteps, mean_reward, std_reward))
|
|
||||||
print("FPS: {:.2f}".format(self.num_timesteps / (time.time() - self.start_time)))
|
|
||||||
|
|
||||||
actor_steps = 0
|
actor_steps = 0
|
||||||
# evaluate all actors
|
# evaluate all actors
|
||||||
for params in self.es_params:
|
for params in self.es_params:
|
||||||
|
|
@ -180,7 +166,6 @@ class CEMRL(TD3):
|
||||||
n_steps=-1, action_noise=self.action_noise,
|
n_steps=-1, action_noise=self.action_noise,
|
||||||
deterministic=False, callback=callback,
|
deterministic=False, callback=callback,
|
||||||
learning_starts=self.learning_starts,
|
learning_starts=self.learning_starts,
|
||||||
num_timesteps=self.num_timesteps,
|
|
||||||
replay_buffer=self.replay_buffer,
|
replay_buffer=self.replay_buffer,
|
||||||
obs=obs, episode_num=episode_num,
|
obs=obs, episode_num=episode_num,
|
||||||
log_interval=log_interval)
|
log_interval=log_interval)
|
||||||
|
|
@ -192,8 +177,6 @@ class CEMRL(TD3):
|
||||||
break
|
break
|
||||||
|
|
||||||
episode_num += n_episodes
|
episode_num += n_episodes
|
||||||
self.num_timesteps += episode_timesteps
|
|
||||||
timesteps_since_eval += episode_timesteps
|
|
||||||
actor_steps += episode_timesteps
|
actor_steps += episode_timesteps
|
||||||
self.fitnesses.append(episode_reward)
|
self.fitnesses.append(episode_reward)
|
||||||
|
|
||||||
|
|
@ -202,7 +185,6 @@ class CEMRL(TD3):
|
||||||
|
|
||||||
self._update_current_progress(self.num_timesteps, total_timesteps)
|
self._update_current_progress(self.num_timesteps, total_timesteps)
|
||||||
self.es.tell(self.es_params, self.fitnesses)
|
self.es.tell(self.es_params, self.fitnesses)
|
||||||
timesteps_since_eval += actor_steps
|
|
||||||
|
|
||||||
callback.on_training_end()
|
callback.on_training_end()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,8 +2,7 @@ import time
|
||||||
import os
|
import os
|
||||||
import io
|
import io
|
||||||
import zipfile
|
import zipfile
|
||||||
import typing
|
from typing import Union, Type, Optional, Dict, Any, List, Tuple, Callable
|
||||||
from typing import Union, Type, Optional, Dict, Any, List, Tuple
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from collections import deque
|
from collections import deque
|
||||||
|
|
||||||
|
|
@ -14,16 +13,13 @@ import numpy as np
|
||||||
from torchy_baselines.common import logger
|
from torchy_baselines.common import logger
|
||||||
from torchy_baselines.common.policies import BasePolicy, get_policy_from_name
|
from torchy_baselines.common.policies import BasePolicy, get_policy_from_name
|
||||||
from torchy_baselines.common.utils import set_random_seed, get_schedule_fn, update_learning_rate
|
from torchy_baselines.common.utils import set_random_seed, get_schedule_fn, update_learning_rate
|
||||||
from torchy_baselines.common.vec_env import DummyVecEnv, VecEnv, unwrap_vec_normalize, sync_envs_normalization
|
from torchy_baselines.common.vec_env import DummyVecEnv, VecEnv, unwrap_vec_normalize
|
||||||
from torchy_baselines.common.monitor import Monitor
|
from torchy_baselines.common.monitor import Monitor
|
||||||
from torchy_baselines.common.evaluation import evaluate_policy
|
|
||||||
from torchy_baselines.common.save_util import data_to_json, json_to_data
|
from torchy_baselines.common.save_util import data_to_json, json_to_data
|
||||||
from torchy_baselines.common.type_aliases import GymEnv, TensorDict, OptimizerStateDict
|
from torchy_baselines.common.type_aliases import GymEnv, TensorDict, OptimizerStateDict
|
||||||
|
from torchy_baselines.common.callbacks import BaseCallback, CallbackList, ConvertCallback, EvalCallback
|
||||||
from torchy_baselines.common.noise import ActionNoise
|
from torchy_baselines.common.noise import ActionNoise
|
||||||
|
|
||||||
if typing.TYPE_CHECKING:
|
|
||||||
from torchy_baselines.common.callbacks import BaseCallback
|
|
||||||
|
|
||||||
|
|
||||||
class BaseRLModel(ABC):
|
class BaseRLModel(ABC):
|
||||||
"""
|
"""
|
||||||
|
|
@ -50,11 +46,12 @@ class BaseRLModel(ABC):
|
||||||
:param sde_sample_freq: Sample a new noise matrix every n steps when using SDE
|
:param sde_sample_freq: Sample a new noise matrix every n steps when using SDE
|
||||||
Default: -1 (only sample at the beginning of the rollout)
|
Default: -1 (only sample at the beginning of the rollout)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self,
|
def __init__(self,
|
||||||
policy: Type[BasePolicy],
|
policy: Type[BasePolicy],
|
||||||
env: Union[GymEnv, str],
|
env: Union[GymEnv, str],
|
||||||
policy_base: Type[BasePolicy],
|
policy_base: Type[BasePolicy],
|
||||||
policy_kwargs : Dict[str, Any] = None,
|
policy_kwargs: Dict[str, Any] = None,
|
||||||
verbose: int = 0,
|
verbose: int = 0,
|
||||||
device: Union[th.device, str] = 'auto',
|
device: Union[th.device, str] = 'auto',
|
||||||
support_multi_env: bool = False,
|
support_multi_env: bool = False,
|
||||||
|
|
@ -133,9 +130,6 @@ class BaseRLModel(ABC):
|
||||||
def _get_eval_env(self, eval_env: Optional[GymEnv]) -> Optional[GymEnv]:
|
def _get_eval_env(self, eval_env: Optional[GymEnv]) -> Optional[GymEnv]:
|
||||||
"""
|
"""
|
||||||
Return the environment that will be used for evaluation.
|
Return the environment that will be used for evaluation.
|
||||||
|
|
||||||
:param eval_env:
|
|
||||||
:return:
|
|
||||||
"""
|
"""
|
||||||
if eval_env is None:
|
if eval_env is None:
|
||||||
eval_env = self.eval_env
|
eval_env = self.eval_env
|
||||||
|
|
@ -146,34 +140,10 @@ class BaseRLModel(ABC):
|
||||||
assert eval_env.num_envs == 1
|
assert eval_env.num_envs == 1
|
||||||
return eval_env
|
return eval_env
|
||||||
|
|
||||||
# Type hint as string to avoid circular import
|
|
||||||
def _init_callback(self, callback) -> 'BaseCallback':
|
|
||||||
"""
|
|
||||||
Note: we cannot use type hint here because of circular import.
|
|
||||||
|
|
||||||
:param callback: (Union[callable, [BaseCallback], BaseCallback, None])
|
|
||||||
:return: (BaseCallback)
|
|
||||||
"""
|
|
||||||
# Avoid circular import
|
|
||||||
from torchy_baselines.common.callbacks import BaseCallback, CallbackList, ConvertCallback
|
|
||||||
|
|
||||||
# Convert a list of callbacks into a callback
|
|
||||||
if isinstance(callback, list):
|
|
||||||
callback = CallbackList(callback)
|
|
||||||
# Convert functional callback to object
|
|
||||||
if not isinstance(callback, BaseCallback):
|
|
||||||
callback = ConvertCallback(callback)
|
|
||||||
|
|
||||||
callback.init_callback(self)
|
|
||||||
return callback
|
|
||||||
|
|
||||||
def scale_action(self, action: np.ndarray) -> np.ndarray:
|
def scale_action(self, action: np.ndarray) -> np.ndarray:
|
||||||
"""
|
"""
|
||||||
Rescale the action from [low, high] to [-1, 1]
|
Rescale the action from [low, high] to [-1, 1]
|
||||||
(no need for symmetric action space)
|
(no need for symmetric action space)
|
||||||
|
|
||||||
:param action:
|
|
||||||
:return:
|
|
||||||
"""
|
"""
|
||||||
low, high = self.action_space.low, self.action_space.high
|
low, high = self.action_space.low, self.action_space.high
|
||||||
return 2.0 * ((action - low) / (high - low)) - 1.0
|
return 2.0 * ((action - low) / (high - low)) - 1.0
|
||||||
|
|
@ -182,9 +152,6 @@ class BaseRLModel(ABC):
|
||||||
"""
|
"""
|
||||||
Rescale the action from [-1, 1] to [low, high]
|
Rescale the action from [-1, 1] to [low, high]
|
||||||
(no need for symmetric action space)
|
(no need for symmetric action space)
|
||||||
|
|
||||||
:param scaled_action:
|
|
||||||
:return:
|
|
||||||
"""
|
"""
|
||||||
low, high = self.action_space.low, self.action_space.high
|
low, high = self.action_space.low, self.action_space.high
|
||||||
return low + (0.5 * (scaled_action + 1.0) * (high - low))
|
return low + (0.5 * (scaled_action + 1.0) * (high - low))
|
||||||
|
|
@ -291,7 +258,7 @@ class BaseRLModel(ABC):
|
||||||
return self.policy.state_dict()
|
return self.policy.state_dict()
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def get_opt_parameters(self)-> OptimizerStateDict:
|
def get_opt_parameters(self) -> OptimizerStateDict:
|
||||||
"""
|
"""
|
||||||
Get current model optimizer parameters as dictionary of variable names -> tensors
|
Get current model optimizer parameters as dictionary of variable names -> tensors
|
||||||
:return: (dict) Dictionary of variable name -> tensor of model's optimizer parameters
|
:return: (dict) Dictionary of variable name -> tensor of model's optimizer parameters
|
||||||
|
|
@ -300,11 +267,13 @@ class BaseRLModel(ABC):
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def learn(self, total_timesteps: int,
|
def learn(self, total_timesteps: int,
|
||||||
callback=None, log_interval: int = 100,
|
callback: Union[None, Callable, List[BaseCallback], BaseCallback] = None,
|
||||||
|
log_interval: int = 100,
|
||||||
tb_log_name: str = "run",
|
tb_log_name: str = "run",
|
||||||
eval_env: Optional[GymEnv] = None,
|
eval_env: Optional[GymEnv] = None,
|
||||||
eval_freq: int = -1,
|
eval_freq: int = -1,
|
||||||
n_eval_episodes: int = 5,
|
n_eval_episodes: int = 5,
|
||||||
|
eval_log_path: Optional[str] = None,
|
||||||
reset_num_timesteps: bool = True):
|
reset_num_timesteps: bool = True):
|
||||||
"""
|
"""
|
||||||
Return a trained model.
|
Return a trained model.
|
||||||
|
|
@ -318,6 +287,8 @@ class BaseRLModel(ABC):
|
||||||
:param eval_env: (gym.Env) Environment that will be used to evaluate the agent
|
:param eval_env: (gym.Env) Environment that will be used to evaluate the agent
|
||||||
:param eval_freq: (int) Evaluate the agent every `eval_freq` timesteps (this may vary a little)
|
:param eval_freq: (int) Evaluate the agent every `eval_freq` timesteps (this may vary a little)
|
||||||
:param n_eval_episodes: (int) Number of episode to evaluate the agent
|
:param n_eval_episodes: (int) Number of episode to evaluate the agent
|
||||||
|
:param eval_log_path: (Optional[str]) Path to a folder where the evaluations will be saved
|
||||||
|
:param reset_num_timesteps: (bool)
|
||||||
:return: (BaseRLModel) the trained model
|
:return: (BaseRLModel) the trained model
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
|
|
@ -391,7 +362,8 @@ class BaseRLModel(ABC):
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _load_from_file(load_path: str, load_data: bool = True) -> (Tuple[Optional[Dict[str, Any]],
|
def _load_from_file(load_path: str, load_data: bool = True) -> (Tuple[Optional[Dict[str, Any]],
|
||||||
Optional[TensorDict], Optional[OptimizerStateDict]]):
|
Optional[TensorDict],
|
||||||
|
Optional[OptimizerStateDict]]):
|
||||||
""" Load model data from a .zip archive
|
""" Load model data from a .zip archive
|
||||||
|
|
||||||
:param load_path: Where to load the model from
|
:param load_path: Where to load the model from
|
||||||
|
|
@ -473,14 +445,53 @@ class BaseRLModel(ABC):
|
||||||
if self.eval_env is not None:
|
if self.eval_env is not None:
|
||||||
self.eval_env.seed(seed)
|
self.eval_env.seed(seed)
|
||||||
|
|
||||||
def _setup_learn(self, eval_env: Optional[GymEnv], callback=None) -> (Tuple[int, int,
|
def _init_callback(self,
|
||||||
List[Any], np.ndarray, Optional[VecEnv], Any]):
|
callback: Union[None, Callable, List[BaseCallback], BaseCallback],
|
||||||
|
eval_env: Optional[VecEnv] = None,
|
||||||
|
eval_freq: int = 10000,
|
||||||
|
n_eval_episodes: int = 5,
|
||||||
|
log_path: Optional[str] = None) -> BaseCallback:
|
||||||
|
"""
|
||||||
|
:param callback: (Union[callable, [BaseCallback], BaseCallback, None])
|
||||||
|
:return: (BaseCallback)
|
||||||
|
"""
|
||||||
|
# Convert a list of callbacks into a callback
|
||||||
|
if isinstance(callback, list):
|
||||||
|
callback = CallbackList(callback)
|
||||||
|
|
||||||
|
# Convert functional callback to object
|
||||||
|
if not isinstance(callback, BaseCallback):
|
||||||
|
callback = ConvertCallback(callback)
|
||||||
|
|
||||||
|
# Create eval callback in charge of the evaluation
|
||||||
|
if eval_env is not None:
|
||||||
|
# Same folder as the rest
|
||||||
|
best_model_save_path = os.path.dirname(log_path) if log_path is not None else None
|
||||||
|
|
||||||
|
eval_callback = EvalCallback(eval_env,
|
||||||
|
best_model_save_path=best_model_save_path,
|
||||||
|
log_path=log_path, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes)
|
||||||
|
callback = CallbackList([callback, eval_callback])
|
||||||
|
|
||||||
|
callback.init_callback(self)
|
||||||
|
return callback
|
||||||
|
|
||||||
|
def _setup_learn(self,
|
||||||
|
eval_env: Optional[GymEnv],
|
||||||
|
callback: Union[None, Callable, List[BaseCallback], BaseCallback] = None,
|
||||||
|
eval_freq: int = 10000,
|
||||||
|
n_eval_episodes: int = 5,
|
||||||
|
log_path: Optional[str] = None
|
||||||
|
) -> Tuple[int, np.ndarray, BaseCallback]:
|
||||||
"""
|
"""
|
||||||
Initialize different variables needed for training.
|
Initialize different variables needed for training.
|
||||||
|
|
||||||
:param eval_env: (Optional[GymEnv])
|
:param eval_env: (Optional[GymEnv])
|
||||||
:param callback: (Union[None, BaseCallback, List[BaseCallback, Callable]])
|
:param callback: (Union[None, BaseCallback, List[BaseCallback, Callable]])
|
||||||
:return: (int, int, [float], np.ndarray, VecEnv, BaseCallback)
|
:param eval_freq: (int)
|
||||||
|
:param n_eval_episodes: (int)
|
||||||
|
:param log_path (Optional[str]):
|
||||||
|
:return: (Tuple[int, np.ndarray, BaseCallback])
|
||||||
"""
|
"""
|
||||||
self.start_time = time.time()
|
self.start_time = time.time()
|
||||||
self.ep_info_buffer = deque(maxlen=100)
|
self.ep_info_buffer = deque(maxlen=100)
|
||||||
|
|
@ -488,18 +499,18 @@ class BaseRLModel(ABC):
|
||||||
if self.action_noise is not None:
|
if self.action_noise is not None:
|
||||||
self.action_noise.reset()
|
self.action_noise.reset()
|
||||||
|
|
||||||
callback = self._init_callback(callback)
|
|
||||||
|
|
||||||
timesteps_since_eval, episode_num = 0, 0
|
timesteps_since_eval, episode_num = 0, 0
|
||||||
evaluations = []
|
|
||||||
|
|
||||||
if eval_env is not None and self.seed is not None:
|
if eval_env is not None and self.seed is not None:
|
||||||
eval_env.seed(self.seed)
|
eval_env.seed(self.seed)
|
||||||
|
|
||||||
eval_env = self._get_eval_env(eval_env)
|
eval_env = self._get_eval_env(eval_env)
|
||||||
obs = self.env.reset() # type: GymEnv
|
obs = self.env.reset()
|
||||||
|
|
||||||
return timesteps_since_eval, episode_num, evaluations, obs, eval_env, callback
|
# Create eval callback if needed
|
||||||
|
callback = self._init_callback(callback, eval_env, eval_freq, n_eval_episodes, log_path)
|
||||||
|
|
||||||
|
return episode_num, obs, callback
|
||||||
|
|
||||||
def _update_info_buffer(self, infos: List[Dict[str, Any]]) -> None:
|
def _update_info_buffer(self, infos: List[Dict[str, Any]]) -> None:
|
||||||
"""
|
"""
|
||||||
|
|
@ -521,7 +532,6 @@ class BaseRLModel(ABC):
|
||||||
action_noise: Optional[ActionNoise] = None,
|
action_noise: Optional[ActionNoise] = None,
|
||||||
deterministic: bool = False,
|
deterministic: bool = False,
|
||||||
learning_starts: int = 0,
|
learning_starts: int = 0,
|
||||||
num_timesteps: int = 0,
|
|
||||||
replay_buffer=None,
|
replay_buffer=None,
|
||||||
obs: Optional[np.ndarray] = None,
|
obs: Optional[np.ndarray] = None,
|
||||||
episode_num: int = 0,
|
episode_num: int = 0,
|
||||||
|
|
@ -537,7 +547,6 @@ class BaseRLModel(ABC):
|
||||||
:param deterministic: (bool)
|
:param deterministic: (bool)
|
||||||
:param callback: (BaseCallback)
|
:param callback: (BaseCallback)
|
||||||
:param learning_starts: (int)
|
:param learning_starts: (int)
|
||||||
:param num_timesteps: (int)
|
|
||||||
:param replay_buffer: (ReplayBuffer)
|
:param replay_buffer: (ReplayBuffer)
|
||||||
:param obs: (np.ndarray)
|
:param obs: (np.ndarray)
|
||||||
:param episode_num: (int)
|
:param episode_num: (int)
|
||||||
|
|
@ -583,7 +592,7 @@ class BaseRLModel(ABC):
|
||||||
# Select action randomly or according to policy
|
# Select action randomly or according to policy
|
||||||
# TODO: use action from policy when using SDE during the warmup phase?
|
# TODO: use action from policy when using SDE during the warmup phase?
|
||||||
# if num_timesteps < learning_starts and not self.use_sde:
|
# if num_timesteps < learning_starts and not self.use_sde:
|
||||||
if num_timesteps < learning_starts:
|
if self.num_timesteps < learning_starts:
|
||||||
# Warmup phase
|
# Warmup phase
|
||||||
unscaled_action = np.array([self.action_space.sample()])
|
unscaled_action = np.array([self.action_space.sample()])
|
||||||
else:
|
else:
|
||||||
|
|
@ -642,7 +651,7 @@ class BaseRLModel(ABC):
|
||||||
if self._vec_normalize_env is not None:
|
if self._vec_normalize_env is not None:
|
||||||
obs_ = new_obs_
|
obs_ = new_obs_
|
||||||
|
|
||||||
num_timesteps += 1
|
self.num_timesteps += 1
|
||||||
episode_timesteps += 1
|
episode_timesteps += 1
|
||||||
total_steps += 1
|
total_steps += 1
|
||||||
if 0 < n_steps <= total_steps:
|
if 0 < n_steps <= total_steps:
|
||||||
|
|
@ -658,8 +667,8 @@ class BaseRLModel(ABC):
|
||||||
|
|
||||||
# Display training infos
|
# Display training infos
|
||||||
if self.verbose >= 1 and log_interval is not None and (
|
if self.verbose >= 1 and log_interval is not None and (
|
||||||
episode_num + total_episodes) % log_interval == 0:
|
episode_num + total_episodes) % log_interval == 0:
|
||||||
fps = int(num_timesteps / (time.time() - self.start_time))
|
fps = int(self.num_timesteps / (time.time() - self.start_time))
|
||||||
logger.logkv("episodes", episode_num + total_episodes)
|
logger.logkv("episodes", episode_num + total_episodes)
|
||||||
if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0:
|
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_rew_mean', self.safe_mean([ep_info['r'] for ep_info in self.ep_info_buffer]))
|
||||||
|
|
@ -667,7 +676,7 @@ class BaseRLModel(ABC):
|
||||||
# logger.logkv("n_updates", n_updates)
|
# logger.logkv("n_updates", n_updates)
|
||||||
logger.logkv("fps", fps)
|
logger.logkv("fps", fps)
|
||||||
logger.logkv('time_elapsed', int(time.time() - self.start_time))
|
logger.logkv('time_elapsed', int(time.time() - self.start_time))
|
||||||
logger.logkv("total timesteps", num_timesteps)
|
logger.logkv("total timesteps", self.num_timesteps)
|
||||||
if self.use_sde:
|
if self.use_sde:
|
||||||
logger.logkv("std", (self.actor.get_std()).mean().item())
|
logger.logkv("std", (self.actor.get_std()).mean().item())
|
||||||
logger.dumpkvs()
|
logger.dumpkvs()
|
||||||
|
|
@ -701,7 +710,7 @@ class BaseRLModel(ABC):
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _save_to_file_zip(save_path: str, data: Dict[str, Any] = None,
|
def _save_to_file_zip(save_path: str, data: Dict[str, Any] = None,
|
||||||
params: TensorDict = None, opt_params: OptimizerStateDict = None) -> None:
|
params: TensorDict = None, opt_params: OptimizerStateDict = None) -> None:
|
||||||
"""
|
"""
|
||||||
Save model to a zip archive.
|
Save model to a zip archive.
|
||||||
|
|
||||||
|
|
@ -776,27 +785,3 @@ class BaseRLModel(ABC):
|
||||||
params_to_save = self.get_policy_parameters()
|
params_to_save = self.get_policy_parameters()
|
||||||
opt_params_to_save = self.get_opt_parameters()
|
opt_params_to_save = self.get_opt_parameters()
|
||||||
self._save_to_file_zip(path, data=data, params=params_to_save, opt_params=opt_params_to_save)
|
self._save_to_file_zip(path, data=data, params=params_to_save, opt_params=opt_params_to_save)
|
||||||
|
|
||||||
def _eval_policy(self, eval_freq: int, eval_env: GymEnv, n_eval_episodes: int,
|
|
||||||
timesteps_since_eval: int, render: bool = False, deterministic: bool = True) -> int:
|
|
||||||
"""
|
|
||||||
Evaluate the current policy on a test environment.
|
|
||||||
|
|
||||||
:param eval_freq: Evaluate the agent every `eval_freq` timesteps (this may vary a little)
|
|
||||||
:param n_eval_episodes: Number of episode to evaluate the agent
|
|
||||||
:parma timesteps_since_eval: Number of timesteps since last evaluation
|
|
||||||
:param deterministic: Whether to use deterministic or stochastic actions
|
|
||||||
:param render: Whether to render the eval env or not
|
|
||||||
:return: Number of timesteps since last evaluation
|
|
||||||
"""
|
|
||||||
if 0 < eval_freq <= timesteps_since_eval and eval_env is not None:
|
|
||||||
timesteps_since_eval %= eval_freq
|
|
||||||
# Synchronise the normalization stats if needed
|
|
||||||
sync_envs_normalization(self.env, eval_env)
|
|
||||||
mean_reward, std_reward = evaluate_policy(self, eval_env, n_eval_episodes,
|
|
||||||
render=render, deterministic=deterministic)
|
|
||||||
if self.verbose > 0:
|
|
||||||
print(f"Eval num_timesteps={self.num_timesteps}, "
|
|
||||||
f"episode_reward={mean_reward:.2f} +/- {std_reward:.2f}")
|
|
||||||
print(f"FPS: {self.num_timesteps / (time.time() - self.start_time):.2f}")
|
|
||||||
return timesteps_since_eval
|
|
||||||
|
|
|
||||||
|
|
@ -1,15 +1,18 @@
|
||||||
import os
|
import os
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
|
import typing
|
||||||
from typing import Union, List, Dict, Any, Optional
|
from typing import Union, List, Dict, Any, Optional
|
||||||
|
|
||||||
import gym
|
import gym
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from torchy_baselines.common.base_class import BaseRLModel # pytype: disable=pyi-error
|
|
||||||
from torchy_baselines.common.vec_env import VecEnv, sync_envs_normalization
|
from torchy_baselines.common.vec_env import VecEnv, sync_envs_normalization
|
||||||
from torchy_baselines.common.evaluation import evaluate_policy
|
from torchy_baselines.common.evaluation import evaluate_policy
|
||||||
from torchy_baselines.common.logger import Logger
|
from torchy_baselines.common.logger import Logger
|
||||||
|
|
||||||
|
if typing.TYPE_CHECKING:
|
||||||
|
from torchy_baselines.common.base_class import BaseRLModel # pytype: disable=pyi-error
|
||||||
|
|
||||||
|
|
||||||
class BaseCallback(ABC):
|
class BaseCallback(ABC):
|
||||||
"""
|
"""
|
||||||
|
|
@ -31,7 +34,8 @@ class BaseCallback(ABC):
|
||||||
# to have access to the parent object
|
# to have access to the parent object
|
||||||
self.parent = None # type: Optional[BaseCallback]
|
self.parent = None # type: Optional[BaseCallback]
|
||||||
|
|
||||||
def init_callback(self, model: BaseRLModel) -> None:
|
# Type hint as string to avoid circular import
|
||||||
|
def init_callback(self, model: 'BaseRLModel') -> None:
|
||||||
"""
|
"""
|
||||||
Initialize the callback by saving references to the
|
Initialize the callback by saving references to the
|
||||||
RL model and the training environment for convenience.
|
RL model and the training environment for convenience.
|
||||||
|
|
@ -105,18 +109,23 @@ class EventCallback(BaseCallback):
|
||||||
if callback is not None:
|
if callback is not None:
|
||||||
self.callback.parent = self
|
self.callback.parent = self
|
||||||
|
|
||||||
def init_callback(self, model: BaseRLModel) -> None:
|
def init_callback(self, model: 'BaseRLModel') -> None:
|
||||||
super(EventCallback, self).init_callback(model)
|
super(EventCallback, self).init_callback(model)
|
||||||
self.callback.init_callback(self.model)
|
if self.callback is not None:
|
||||||
|
self.callback.init_callback(self.model)
|
||||||
|
|
||||||
def _on_training_start(self) -> None:
|
def _on_training_start(self) -> None:
|
||||||
self.callback.on_training_start(self.locals, self.globals)
|
if self.callback is not None:
|
||||||
|
self.callback.on_training_start(self.locals, self.globals)
|
||||||
|
|
||||||
def _on_event(self) -> bool:
|
def _on_event(self) -> bool:
|
||||||
if self.callback is not None:
|
if self.callback is not None:
|
||||||
return self.callback()
|
return self.callback()
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
def _on_step(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
class CallbackList(BaseCallback):
|
class CallbackList(BaseCallback):
|
||||||
def __init__(self, callbacks: List[BaseCallback]):
|
def __init__(self, callbacks: List[BaseCallback]):
|
||||||
|
|
@ -179,7 +188,7 @@ class ConvertCallback(BaseCallback):
|
||||||
"""
|
"""
|
||||||
Convert functional callback (old-style) to object.
|
Convert functional callback (old-style) to object.
|
||||||
|
|
||||||
:param on_step: (callable)
|
:param callback: (callable)
|
||||||
:param verbose: (int)
|
:param verbose: (int)
|
||||||
"""
|
"""
|
||||||
def __init__(self, callback, verbose=0):
|
def __init__(self, callback, verbose=0):
|
||||||
|
|
@ -207,6 +216,7 @@ class EvalCallback(EventCallback):
|
||||||
according to performance on the eval env will be saved.
|
according to performance on the eval env will be saved.
|
||||||
:param deterministic: (bool) Whether the evaluation should
|
:param deterministic: (bool) Whether the evaluation should
|
||||||
use a stochastic or deterministic actions.
|
use a stochastic or deterministic actions.
|
||||||
|
:param deterministic: (bool) Whether to render or not the environment during evaluation
|
||||||
:param verbose: (int)
|
:param verbose: (int)
|
||||||
"""
|
"""
|
||||||
def __init__(self, eval_env: Union[gym.Env, VecEnv],
|
def __init__(self, eval_env: Union[gym.Env, VecEnv],
|
||||||
|
|
@ -216,12 +226,15 @@ class EvalCallback(EventCallback):
|
||||||
log_path: str = None,
|
log_path: str = None,
|
||||||
best_model_save_path: str = None,
|
best_model_save_path: str = None,
|
||||||
deterministic: bool = True,
|
deterministic: bool = True,
|
||||||
|
render: bool = False,
|
||||||
verbose: int = 1):
|
verbose: int = 1):
|
||||||
super(EvalCallback, self).__init__(callback_on_new_best, verbose=verbose)
|
super(EvalCallback, self).__init__(callback_on_new_best, verbose=verbose)
|
||||||
self.n_eval_episodes = n_eval_episodes
|
self.n_eval_episodes = n_eval_episodes
|
||||||
self.eval_freq = eval_freq
|
self.eval_freq = eval_freq
|
||||||
self.best_mean_reward = -np.inf
|
self.best_mean_reward = -np.inf
|
||||||
self.deterministic = deterministic
|
self.deterministic = deterministic
|
||||||
|
self.render = render
|
||||||
|
|
||||||
if isinstance(eval_env, VecEnv):
|
if isinstance(eval_env, VecEnv):
|
||||||
assert eval_env.num_envs == 1, "You must pass only one environment for evaluation"
|
assert eval_env.num_envs == 1, "You must pass only one environment for evaluation"
|
||||||
|
|
||||||
|
|
@ -230,6 +243,7 @@ class EvalCallback(EventCallback):
|
||||||
self.log_path = log_path
|
self.log_path = log_path
|
||||||
self.evaluations_results = []
|
self.evaluations_results = []
|
||||||
self.evaluations_timesteps = []
|
self.evaluations_timesteps = []
|
||||||
|
self.evaluations_length = []
|
||||||
|
|
||||||
def _init_callback(self):
|
def _init_callback(self):
|
||||||
# Does not work when eval_env is a gym.Env and training_env is a VecEnv
|
# Does not work when eval_env is a gym.Env and training_env is a VecEnv
|
||||||
|
|
@ -244,22 +258,30 @@ class EvalCallback(EventCallback):
|
||||||
|
|
||||||
def _on_step(self) -> bool:
|
def _on_step(self) -> bool:
|
||||||
|
|
||||||
if self.n_calls % self.eval_freq == 0:
|
if self.eval_freq > 0 and self.n_calls % self.eval_freq == 0:
|
||||||
# Sync training and eval env if there is VecNormalize
|
# Sync training and eval env if there is VecNormalize
|
||||||
sync_envs_normalization(self.training_env, self.eval_env)
|
sync_envs_normalization(self.training_env, self.eval_env)
|
||||||
|
|
||||||
episode_rewards, _ = evaluate_policy(self.model, self.eval_env, n_eval_episodes=self.n_eval_episodes,
|
episode_rewards, episode_lengths = evaluate_policy(self.model, self.eval_env,
|
||||||
deterministic=self.deterministic, return_episode_rewards=True)
|
n_eval_episodes=self.n_eval_episodes,
|
||||||
|
render=self.render,
|
||||||
|
deterministic=self.deterministic,
|
||||||
|
return_episode_rewards=True)
|
||||||
|
|
||||||
if self.log_path is not None:
|
if self.log_path is not None:
|
||||||
self.evaluations_timesteps.append(self.num_timesteps)
|
self.evaluations_timesteps.append(self.num_timesteps)
|
||||||
self.evaluations_results.append(episode_rewards)
|
self.evaluations_results.append(episode_rewards)
|
||||||
np.savez(self.log_path, timesteps=self.evaluations_timesteps, results=self.evaluations_results)
|
self.evaluations_length.append(episode_lengths)
|
||||||
|
np.savez(self.log_path, timesteps=self.evaluations_timesteps,
|
||||||
|
results=self.evaluations_results, ep_lengths=self.evaluations_length)
|
||||||
|
|
||||||
mean_reward, std_reward = np.mean(episode_rewards), np.std(episode_rewards)
|
mean_reward, std_reward = np.mean(episode_rewards), np.std(episode_rewards)
|
||||||
|
mean_ep_length, std_ep_length = np.mean(episode_lengths), np.std(episode_lengths)
|
||||||
|
|
||||||
if self.verbose > 0:
|
if self.verbose > 0:
|
||||||
print(f"Eval num_timesteps={self.num_timesteps}, "
|
print(f"Eval num_timesteps={self.num_timesteps}, "
|
||||||
f"episode_reward={mean_reward:.2f} +/- {std_reward:.2f}")
|
f"episode_reward={mean_reward:.2f} +/- {std_reward:.2f}")
|
||||||
|
print(f"Episode length: {mean_ep_length:.2f} +/- {std_ep_length:.2f}")
|
||||||
|
|
||||||
if mean_reward > self.best_mean_reward:
|
if mean_reward > self.best_mean_reward:
|
||||||
if self.verbose > 0:
|
if self.verbose > 0:
|
||||||
|
|
|
||||||
|
|
@ -24,31 +24,33 @@ def evaluate_policy(model, env, n_eval_episodes=10, deterministic=True,
|
||||||
:param return_episode_rewards: (bool) If True, a list of reward per episode
|
:param return_episode_rewards: (bool) If True, a list of reward per episode
|
||||||
will be returned instead of the mean.
|
will be returned instead of the mean.
|
||||||
:return: (float, float) Mean reward per episode, std of reward per episode
|
:return: (float, float) Mean reward per episode, std of reward per episode
|
||||||
returns ([float], int) when `return_episode_rewards` is True
|
returns ([float], [int]) when `return_episode_rewards` is True
|
||||||
"""
|
"""
|
||||||
if isinstance(env, VecEnv):
|
if isinstance(env, VecEnv):
|
||||||
assert env.num_envs == 1, "You must pass only one environment when using this function"
|
assert env.num_envs == 1, "You must pass only one environment when using this function"
|
||||||
|
|
||||||
episode_rewards, n_steps = [], 0
|
episode_rewards, episode_lengths = [], []
|
||||||
for _ in range(n_eval_episodes):
|
for _ in range(n_eval_episodes):
|
||||||
obs = env.reset()
|
obs = env.reset()
|
||||||
done = False
|
done = False
|
||||||
episode_reward = 0.0
|
episode_reward = 0.0
|
||||||
|
episode_length = 0
|
||||||
while not done:
|
while not done:
|
||||||
action = model.predict(obs, deterministic=deterministic)
|
action = model.predict(obs, deterministic=deterministic)
|
||||||
obs, reward, done, _info = env.step(action)
|
obs, reward, done, _info = env.step(action)
|
||||||
episode_reward += reward
|
episode_reward += reward
|
||||||
if callback is not None:
|
if callback is not None:
|
||||||
callback(locals(), globals())
|
callback(locals(), globals())
|
||||||
n_steps += 1
|
episode_length += 1
|
||||||
if render:
|
if render:
|
||||||
env.render()
|
env.render()
|
||||||
episode_rewards.append(episode_reward)
|
episode_rewards.append(episode_reward)
|
||||||
|
episode_lengths.append(episode_length)
|
||||||
mean_reward = np.mean(episode_rewards)
|
mean_reward = np.mean(episode_rewards)
|
||||||
std_reward = np.std(episode_rewards)
|
std_reward = np.std(episode_rewards)
|
||||||
if reward_threshold is not None:
|
if reward_threshold is not None:
|
||||||
assert mean_reward > reward_threshold, (f'Mean reward below threshold: '
|
assert mean_reward > reward_threshold, (f'Mean reward below threshold: '
|
||||||
'{mean_reward:.2f} < {reward_threshold:.2f}')
|
'{mean_reward:.2f} < {reward_threshold:.2f}')
|
||||||
if return_episode_rewards:
|
if return_episode_rewards:
|
||||||
return episode_rewards, n_steps
|
return episode_rewards, episode_lengths
|
||||||
return mean_reward, std_reward
|
return mean_reward, std_reward
|
||||||
|
|
|
||||||
|
|
@ -193,6 +193,8 @@ class PPO(BaseRLModel):
|
||||||
|
|
||||||
self._update_info_buffer(infos)
|
self._update_info_buffer(infos)
|
||||||
n_steps += 1
|
n_steps += 1
|
||||||
|
self.num_timesteps += env.num_envs
|
||||||
|
|
||||||
if isinstance(self.action_space, gym.spaces.Discrete):
|
if isinstance(self.action_space, gym.spaces.Discrete):
|
||||||
# Reshape in case of discrete action
|
# Reshape in case of discrete action
|
||||||
actions = actions.reshape(-1, 1)
|
actions = actions.reshape(-1, 1)
|
||||||
|
|
@ -284,9 +286,11 @@ class PPO(BaseRLModel):
|
||||||
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
|
||||||
|
|
||||||
def learn(self, total_timesteps, callback=None, log_interval=1,
|
def learn(self, total_timesteps, callback=None, log_interval=1,
|
||||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="PPO", reset_num_timesteps=True):
|
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="PPO",
|
||||||
|
eval_log_path=None, reset_num_timesteps=True):
|
||||||
|
|
||||||
timesteps_since_eval, iteration, evaluations, obs, eval_env, callback = self._setup_learn(eval_env, callback)
|
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq, n_eval_episodes, eval_log_path)
|
||||||
|
iteration = 0
|
||||||
|
|
||||||
if self.tensorboard_log is not None and SummaryWriter is not None:
|
if self.tensorboard_log is not None and SummaryWriter is not None:
|
||||||
self.tb_writer = SummaryWriter(log_dir=os.path.join(self.tensorboard_log, tb_log_name))
|
self.tb_writer = SummaryWriter(log_dir=os.path.join(self.tensorboard_log, tb_log_name))
|
||||||
|
|
@ -295,15 +299,15 @@ class PPO(BaseRLModel):
|
||||||
|
|
||||||
while self.num_timesteps < total_timesteps:
|
while self.num_timesteps < total_timesteps:
|
||||||
|
|
||||||
obs, continue_training = self.collect_rollouts(self.env, callback, self.rollout_buffer, n_rollout_steps=self.n_steps,
|
obs, continue_training = self.collect_rollouts(self.env, callback,
|
||||||
|
self.rollout_buffer,
|
||||||
|
n_rollout_steps=self.n_steps,
|
||||||
obs=obs)
|
obs=obs)
|
||||||
|
|
||||||
if continue_training is False:
|
if continue_training is False:
|
||||||
break
|
break
|
||||||
|
|
||||||
iteration += 1
|
iteration += 1
|
||||||
self.num_timesteps += self.n_steps * self.n_envs
|
|
||||||
timesteps_since_eval += self.n_steps * self.n_envs
|
|
||||||
self._update_current_progress(self.num_timesteps, total_timesteps)
|
self._update_current_progress(self.num_timesteps, total_timesteps)
|
||||||
|
|
||||||
# Display training infos
|
# Display training infos
|
||||||
|
|
@ -320,9 +324,6 @@ class PPO(BaseRLModel):
|
||||||
|
|
||||||
self.train(self.n_epochs, batch_size=self.batch_size)
|
self.train(self.n_epochs, batch_size=self.batch_size)
|
||||||
|
|
||||||
# Evaluate the agent
|
|
||||||
timesteps_since_eval = self._eval_policy(eval_freq, eval_env, n_eval_episodes,
|
|
||||||
timesteps_since_eval, deterministic=True)
|
|
||||||
# For tensorboard integration
|
# For tensorboard integration
|
||||||
# if self.tb_writer is not None:
|
# if self.tb_writer is not None:
|
||||||
# self.tb_writer.add_scalar('Eval/reward', mean_reward, self.num_timesteps)
|
# self.tb_writer.add_scalar('Eval/reward', mean_reward, self.num_timesteps)
|
||||||
|
|
|
||||||
|
|
@ -257,9 +257,9 @@ class SAC(BaseRLModel):
|
||||||
|
|
||||||
def learn(self, total_timesteps, callback=None, log_interval=4,
|
def learn(self, total_timesteps, callback=None, log_interval=4,
|
||||||
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="SAC",
|
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="SAC",
|
||||||
reset_num_timesteps=True):
|
eval_log_path=None, reset_num_timesteps=True):
|
||||||
|
|
||||||
timesteps_since_eval, episode_num, evaluations, obs, eval_env, callback = self._setup_learn(eval_env, callback)
|
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq, n_eval_episodes, eval_log_path)
|
||||||
|
|
||||||
callback.on_training_start(locals(), globals())
|
callback.on_training_start(locals(), globals())
|
||||||
|
|
||||||
|
|
@ -268,7 +268,6 @@ class SAC(BaseRLModel):
|
||||||
n_steps=self.train_freq, action_noise=self.action_noise,
|
n_steps=self.train_freq, action_noise=self.action_noise,
|
||||||
deterministic=False, callback=callback,
|
deterministic=False, callback=callback,
|
||||||
learning_starts=self.learning_starts,
|
learning_starts=self.learning_starts,
|
||||||
num_timesteps=self.num_timesteps,
|
|
||||||
replay_buffer=self.replay_buffer,
|
replay_buffer=self.replay_buffer,
|
||||||
obs=obs, episode_num=episode_num,
|
obs=obs, episode_num=episode_num,
|
||||||
log_interval=log_interval)
|
log_interval=log_interval)
|
||||||
|
|
@ -278,9 +277,7 @@ class SAC(BaseRLModel):
|
||||||
if continue_training is False:
|
if continue_training is False:
|
||||||
break
|
break
|
||||||
|
|
||||||
self.num_timesteps += episode_timesteps
|
|
||||||
episode_num += n_episodes
|
episode_num += n_episodes
|
||||||
timesteps_since_eval += episode_timesteps
|
|
||||||
self._update_current_progress(self.num_timesteps, total_timesteps)
|
self._update_current_progress(self.num_timesteps, total_timesteps)
|
||||||
|
|
||||||
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
|
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
|
||||||
|
|
@ -288,9 +285,6 @@ class SAC(BaseRLModel):
|
||||||
|
|
||||||
self.train(gradient_steps, batch_size=self.batch_size)
|
self.train(gradient_steps, batch_size=self.batch_size)
|
||||||
|
|
||||||
timesteps_since_eval = self._eval_policy(eval_freq, eval_env, n_eval_episodes,
|
|
||||||
timesteps_since_eval, deterministic=True)
|
|
||||||
|
|
||||||
callback.on_training_end()
|
callback.on_training_end()
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -249,9 +249,10 @@ class TD3(BaseRLModel):
|
||||||
del self.rollout_data
|
del self.rollout_data
|
||||||
|
|
||||||
def learn(self, total_timesteps, callback=None, log_interval=4,
|
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):
|
eval_env=None, eval_freq=-1, n_eval_episodes=5,
|
||||||
|
tb_log_name="TD3", eval_log_path=None, reset_num_timesteps=True):
|
||||||
|
|
||||||
timesteps_since_eval, episode_num, evaluations, obs, eval_env, callback = self._setup_learn(eval_env, callback)
|
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq, n_eval_episodes, eval_log_path)
|
||||||
|
|
||||||
callback.on_training_start(locals(), globals())
|
callback.on_training_start(locals(), globals())
|
||||||
|
|
||||||
|
|
@ -261,7 +262,6 @@ class TD3(BaseRLModel):
|
||||||
n_steps=self.train_freq, action_noise=self.action_noise,
|
n_steps=self.train_freq, action_noise=self.action_noise,
|
||||||
deterministic=False, callback=callback,
|
deterministic=False, callback=callback,
|
||||||
learning_starts=self.learning_starts,
|
learning_starts=self.learning_starts,
|
||||||
num_timesteps=self.num_timesteps,
|
|
||||||
replay_buffer=self.replay_buffer,
|
replay_buffer=self.replay_buffer,
|
||||||
obs=obs, episode_num=episode_num,
|
obs=obs, episode_num=episode_num,
|
||||||
log_interval=log_interval)
|
log_interval=log_interval)
|
||||||
|
|
@ -272,8 +272,6 @@ class TD3(BaseRLModel):
|
||||||
break
|
break
|
||||||
|
|
||||||
episode_num += n_episodes
|
episode_num += n_episodes
|
||||||
self.num_timesteps += episode_timesteps
|
|
||||||
timesteps_since_eval += episode_timesteps
|
|
||||||
self._update_current_progress(self.num_timesteps, total_timesteps)
|
self._update_current_progress(self.num_timesteps, total_timesteps)
|
||||||
|
|
||||||
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
|
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
|
||||||
|
|
@ -290,10 +288,6 @@ class TD3(BaseRLModel):
|
||||||
gradient_steps = self.gradient_steps if self.gradient_steps > 0 else episode_timesteps
|
gradient_steps = self.gradient_steps if self.gradient_steps > 0 else episode_timesteps
|
||||||
self.train(gradient_steps, batch_size=self.batch_size, policy_delay=self.policy_delay)
|
self.train(gradient_steps, batch_size=self.batch_size, policy_delay=self.policy_delay)
|
||||||
|
|
||||||
# Evaluate the agent
|
|
||||||
timesteps_since_eval = self._eval_policy(eval_freq, eval_env, n_eval_episodes,
|
|
||||||
timesteps_since_eval, deterministic=True)
|
|
||||||
|
|
||||||
callback.on_training_end()
|
callback.on_training_end()
|
||||||
|
|
||||||
return self
|
return self
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue