From 6710f1576cb3253279db5085cfb0c7021b7a5603 Mon Sep 17 00:00:00 2001 From: Antonin Raffin Date: Fri, 31 Jan 2020 13:48:25 +0100 Subject: [PATCH] Fix eval log path --- torchy_baselines/common/base_class.py | 7 ++----- torchy_baselines/common/callbacks.py | 5 ++++- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/torchy_baselines/common/base_class.py b/torchy_baselines/common/base_class.py index 1f711be..9a3222f 100644 --- a/torchy_baselines/common/base_class.py +++ b/torchy_baselines/common/base_class.py @@ -487,11 +487,8 @@ class BaseRLModel(ABC): # 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, + best_model_save_path=log_path, log_path=log_path, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes) callback = CallbackList([callback, eval_callback]) @@ -513,7 +510,7 @@ class BaseRLModel(ABC): :param callback: (Union[None, BaseCallback, List[BaseCallback, Callable]]) :param eval_freq: (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 :return: (Tuple[int, np.ndarray, BaseCallback]) """ diff --git a/torchy_baselines/common/callbacks.py b/torchy_baselines/common/callbacks.py index 346f115..2481299 100644 --- a/torchy_baselines/common/callbacks.py +++ b/torchy_baselines/common/callbacks.py @@ -240,7 +240,10 @@ class EvalCallback(EventCallback): self.eval_env = eval_env 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_timesteps = [] self.evaluations_length = []