Improve logic for progressbarcallback handling
This commit is contained in:
@@ -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.')
|
||||||
|
|||||||
Reference in New Issue
Block a user