Improve logic for progressbarcallback handling

This commit is contained in:
Matthias
2023-10-15 11:20:25 +02:00
parent ba674fc796
commit 2d9d8dc976
@@ -1,6 +1,6 @@
import logging import logging
from pathlib import Path from pathlib import Path
from typing import Any, Dict, Type from typing import Any, Dict, List, Optional, Type
import torch as th import torch as th
from stable_baselines3.common.callbacks import ProgressBarCallback from stable_baselines3.common.callbacks import ProgressBarCallback
@@ -74,10 +74,11 @@ class ReinforcementLearner(BaseReinforcementLearningModel):
'trained agent.') 'trained agent.')
model = self.dd.model_dictionary[dk.pair] model = self.dd.model_dictionary[dk.pair]
model.set_env(self.train_env) model.set_env(self.train_env)
callbacks = [self.eval_callback, self.tensorboard_callback] callbacks: List[Any] = [self.eval_callback, self.tensorboard_callback]
use_progressbar = self.rl_config.get('progress_bar', False) progressbar_callback: Optional[ProgressBarCallback] = None
if use_progressbar: if self.rl_config.get('progress_bar', False):
callbacks.insert(0, ProgressBarCallback()) progressbar_callback = ProgressBarCallback()
callbacks.insert(0, progressbar_callback)
try: try:
model.learn( model.learn(
@@ -85,8 +86,8 @@ class ReinforcementLearner(BaseReinforcementLearningModel):
callback=callbacks, callback=callbacks,
) )
finally: finally:
if use_progressbar: if progressbar_callback:
callbacks[0].on_training_end() progressbar_callback.on_training_end()
if Path(dk.data_path / "best_model.zip").is_file(): if Path(dk.data_path / "best_model.zip").is_file():
logger.info('Callback found a best model.') logger.info('Callback found a best model.')