Fix eval log path

This commit is contained in:
Antonin Raffin 2020-01-31 13:48:25 +01:00
parent ec657cc34e
commit 6710f1576c
2 changed files with 6 additions and 6 deletions

View file

@ -487,11 +487,8 @@ class BaseRLModel(ABC):
# Create eval callback in charge of the evaluation # Create eval callback in charge of the evaluation
if eval_env is not None: 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, eval_callback = EvalCallback(eval_env,
best_model_save_path=best_model_save_path, best_model_save_path=log_path,
log_path=log_path, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes) log_path=log_path, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes)
callback = CallbackList([callback, eval_callback]) callback = CallbackList([callback, eval_callback])
@ -513,7 +510,7 @@ class BaseRLModel(ABC):
:param callback: (Union[None, BaseCallback, List[BaseCallback, Callable]]) :param callback: (Union[None, BaseCallback, List[BaseCallback, Callable]])
:param eval_freq: (int) :param eval_freq: (int)
:param n_eval_episodes: (int) :param n_eval_episodes: (int)
:param log_path (Optional[str]): :param log_path (Optional[str]): Path to a log folder
:param reset_num_timesteps: (bool) Whether to reset or not the `num_timesteps` attribute :param reset_num_timesteps: (bool) Whether to reset or not the `num_timesteps` attribute
:return: (Tuple[int, np.ndarray, BaseCallback]) :return: (Tuple[int, np.ndarray, BaseCallback])
""" """

View file

@ -240,7 +240,10 @@ class EvalCallback(EventCallback):
self.eval_env = eval_env self.eval_env = eval_env
self.best_model_save_path = best_model_save_path self.best_model_save_path = best_model_save_path
self.log_path = os.path.join(log_path, 'evaluations') # Logs will be written in `evaluations.npz`
if log_path is not None:
os.path.join(log_path, 'evaluations')
self.log_path = log_path
self.evaluations_results = [] self.evaluations_results = []
self.evaluations_timesteps = [] self.evaluations_timesteps = []
self.evaluations_length = [] self.evaluations_length = []