From 8d2b389e2713a82f5f19683d1693a4797fc9e011 Mon Sep 17 00:00:00 2001 From: Matthias Date: Sun, 15 Oct 2023 10:40:45 +0200 Subject: [PATCH 1/7] Fix wording in log msg --- freqtrade/freqai/prediction_models/ReinforcementLearner.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/freqtrade/freqai/prediction_models/ReinforcementLearner.py b/freqtrade/freqai/prediction_models/ReinforcementLearner.py index a11decc92..d9a11a7a8 100644 --- a/freqtrade/freqai/prediction_models/ReinforcementLearner.py +++ b/freqtrade/freqai/prediction_models/ReinforcementLearner.py @@ -85,7 +85,7 @@ class ReinforcementLearner(BaseReinforcementLearningModel): best_model = self.MODELCLASS.load(dk.data_path / "best_model") return best_model - logger.info('Couldnt find best model, using final model instead.') + logger.info("Couldn't find best model, using final model instead.") return model From 646dd63faf89cbad809905349c56c31678435765 Mon Sep 17 00:00:00 2001 From: Matthias Date: Sun, 15 Oct 2023 10:41:07 +0200 Subject: [PATCH 2/7] Properly close out progressbarCallback based on suggestions provided in https://github.com/DLR-RM/stable-baselines3/issues/1645 --- .../prediction_models/ReinforcementLearner.py | 18 +++++++++++++----- 1 file changed, 13 insertions(+), 5 deletions(-) diff --git a/freqtrade/freqai/prediction_models/ReinforcementLearner.py b/freqtrade/freqai/prediction_models/ReinforcementLearner.py index d9a11a7a8..d4c1881a6 100644 --- a/freqtrade/freqai/prediction_models/ReinforcementLearner.py +++ b/freqtrade/freqai/prediction_models/ReinforcementLearner.py @@ -3,6 +3,7 @@ from pathlib import Path from typing import Any, Dict, Type import torch as th +from stable_baselines3.common.callbacks import ProgressBarCallback from freqtrade.freqai.data_kitchen import FreqaiDataKitchen from freqtrade.freqai.RL.Base5ActionRLEnv import Actions, Base5ActionRLEnv, Positions @@ -73,12 +74,19 @@ class ReinforcementLearner(BaseReinforcementLearningModel): 'trained agent.') model = self.dd.model_dictionary[dk.pair] model.set_env(self.train_env) + callbacks = [self.eval_callback, self.tensorboard_callback] + use_progressbar = self.rl_config.get('progress_bar', False) + if use_progressbar: + callbacks.insert(0, ProgressBarCallback()) - model.learn( - total_timesteps=int(total_timesteps), - callback=[self.eval_callback, self.tensorboard_callback], - progress_bar=self.rl_config.get('progress_bar', False) - ) + try: + model.learn( + total_timesteps=int(total_timesteps), + callback=callbacks, + ) + finally: + if use_progressbar: + callbacks[0].on_training_end() if Path(dk.data_path / "best_model.zip").is_file(): logger.info('Callback found a best model.') From 27bae60b68f55e5a38c652e378c46940bf0393b5 Mon Sep 17 00:00:00 2001 From: Matthias Date: Sun, 15 Oct 2023 10:51:36 +0200 Subject: [PATCH 3/7] Fix typo --- freqtrade/freqai/RL/BaseEnvironment.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/freqtrade/freqai/RL/BaseEnvironment.py b/freqtrade/freqai/RL/BaseEnvironment.py index 54502c869..b8548dd16 100644 --- a/freqtrade/freqai/RL/BaseEnvironment.py +++ b/freqtrade/freqai/RL/BaseEnvironment.py @@ -159,7 +159,7 @@ class BaseEnvironment(gym.Env): function is designed for tracking incremented objects, events, actions inside the training environment. For example, a user can call this to track the - frequency of occurence of an `is_valid` call in + frequency of occurrence of an `is_valid` call in their `calculate_reward()`: def calculate_reward(self, action: int) -> float: From 58550515ad8006aac75fd09075620d4490b0ef12 Mon Sep 17 00:00:00 2001 From: Matthias Date: Sun, 15 Oct 2023 11:04:05 +0200 Subject: [PATCH 4/7] Fix deprecation warning from tensorboard-callback --- freqtrade/freqai/tensorboard/TensorboardCallback.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/freqtrade/freqai/tensorboard/TensorboardCallback.py b/freqtrade/freqai/tensorboard/TensorboardCallback.py index 61652c9c6..4cd49689e 100644 --- a/freqtrade/freqai/tensorboard/TensorboardCallback.py +++ b/freqtrade/freqai/tensorboard/TensorboardCallback.py @@ -49,7 +49,7 @@ class TensorboardCallback(BaseCallback): local_info = self.locals["infos"][0] if self.training_env is None: return True - tensorboard_metrics = self.training_env.get_attr("tensorboard_metrics")[0] + tensorboard_metrics = self.training_env.envs[0].unwrapped.tensorboard_metrics for metric in local_info: if metric not in ["episode", "terminal_observation"]: From ba674fc796e32335566e0c1f2c5551fd5115d444 Mon Sep 17 00:00:00 2001 From: Matthias Date: Sun, 15 Oct 2023 11:20:11 +0200 Subject: [PATCH 5/7] Type-ignore training-envs --- freqtrade/freqai/tensorboard/TensorboardCallback.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/freqtrade/freqai/tensorboard/TensorboardCallback.py b/freqtrade/freqai/tensorboard/TensorboardCallback.py index 4cd49689e..aa65aa199 100644 --- a/freqtrade/freqai/tensorboard/TensorboardCallback.py +++ b/freqtrade/freqai/tensorboard/TensorboardCallback.py @@ -49,7 +49,10 @@ class TensorboardCallback(BaseCallback): local_info = self.locals["infos"][0] if self.training_env is None: return True - tensorboard_metrics = self.training_env.envs[0].unwrapped.tensorboard_metrics + + tensorboard_metrics = ( + self.training_env.envs[0].unwrapped.tensorboard_metrics # type: ignore[attr-defined] + ) for metric in local_info: if metric not in ["episode", "terminal_observation"]: From 2d9d8dc976ef1d242420fd9b84d7a8745ea56e7b Mon Sep 17 00:00:00 2001 From: Matthias Date: Sun, 15 Oct 2023 11:20:25 +0200 Subject: [PATCH 6/7] Improve logic for progressbarcallback handling --- .../prediction_models/ReinforcementLearner.py | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/freqtrade/freqai/prediction_models/ReinforcementLearner.py b/freqtrade/freqai/prediction_models/ReinforcementLearner.py index d4c1881a6..fbf12008a 100644 --- a/freqtrade/freqai/prediction_models/ReinforcementLearner.py +++ b/freqtrade/freqai/prediction_models/ReinforcementLearner.py @@ -1,6 +1,6 @@ import logging from pathlib import Path -from typing import Any, Dict, Type +from typing import Any, Dict, List, Optional, Type import torch as th from stable_baselines3.common.callbacks import ProgressBarCallback @@ -74,10 +74,11 @@ class ReinforcementLearner(BaseReinforcementLearningModel): 'trained agent.') model = self.dd.model_dictionary[dk.pair] model.set_env(self.train_env) - callbacks = [self.eval_callback, self.tensorboard_callback] - use_progressbar = self.rl_config.get('progress_bar', False) - if use_progressbar: - callbacks.insert(0, ProgressBarCallback()) + callbacks: List[Any] = [self.eval_callback, self.tensorboard_callback] + progressbar_callback: Optional[ProgressBarCallback] = None + if self.rl_config.get('progress_bar', False): + progressbar_callback = ProgressBarCallback() + callbacks.insert(0, progressbar_callback) try: model.learn( @@ -85,8 +86,8 @@ class ReinforcementLearner(BaseReinforcementLearningModel): callback=callbacks, ) finally: - if use_progressbar: - callbacks[0].on_training_end() + if progressbar_callback: + progressbar_callback.on_training_end() if Path(dk.data_path / "best_model.zip").is_file(): logger.info('Callback found a best model.') From 1e1b8dbe538bba830b6903e1e4b88e63ca8eb6ee Mon Sep 17 00:00:00 2001 From: Matthias Date: Sun, 15 Oct 2023 11:52:18 +0200 Subject: [PATCH 7/7] Handle multiproc calls for now --- freqtrade/freqai/tensorboard/TensorboardCallback.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/freqtrade/freqai/tensorboard/TensorboardCallback.py b/freqtrade/freqai/tensorboard/TensorboardCallback.py index aa65aa199..2be917616 100644 --- a/freqtrade/freqai/tensorboard/TensorboardCallback.py +++ b/freqtrade/freqai/tensorboard/TensorboardCallback.py @@ -50,9 +50,12 @@ class TensorboardCallback(BaseCallback): if self.training_env is None: return True - tensorboard_metrics = ( - self.training_env.envs[0].unwrapped.tensorboard_metrics # type: ignore[attr-defined] - ) + if hasattr(self.training_env, 'envs'): + tensorboard_metrics = self.training_env.envs[0].unwrapped.tensorboard_metrics + + else: + # For RL-multiproc - usage of [0] might need to be evaluated + tensorboard_metrics = self.training_env.get_attr("tensorboard_metrics")[0] for metric in local_info: if metric not in ["episode", "terminal_observation"]: