Source code for callback.hyperparam_callback

from stable_baselines3.common.callbacks import BaseCallback
from stable_baselines3.common.logger import HParam


[docs]class HParamCallback(BaseCallback): """ Saves the hyperparameters and metrics at the start of the training, and logs them to TensorBoard. """ def _on_training_start(self) -> None: hparam_dict = { "algorithm": self.model.__class__.__name__, "learning rate": self.model.learning_rate, "gamma": self.model.gamma, "batch_size": self.model.batch_size, "n_steps": self.model.n_steps } # define the metrics that will appear in the `HPARAMS` Tensorboard tab by referencing their tag # Tensorbaord will find & display metrics from the `SCALARS` tab metric_dict = { "rollout/ep_len_mean": 0, "train/value_loss": 0.0, } self.logger.record( "hparams", HParam(hparam_dict, metric_dict), exclude=("stdout", "log", "json", "csv"), ) def _on_step(self) -> bool: return True