diff --git a/monai/handlers/tensorboard_handlers.py b/monai/handlers/tensorboard_handlers.py index 20e2d74c8c5..9124ec014c6 100644 --- a/monai/handlers/tensorboard_handlers.py +++ b/monai/handlers/tensorboard_handlers.py @@ -27,11 +27,13 @@ from ignite.engine import Engine from tensorboardX import SummaryWriter as SummaryWriterX from torch.utils.tensorboard import SummaryWriter + + _tb_available = True else: Engine, _ = optional_import( "ignite.engine", IgniteInfo.OPT_IMPORT_VERSION, min_version, "Engine", as_type="decorator" ) - SummaryWriter, _ = optional_import("torch.utils.tensorboard", name="SummaryWriter") + SummaryWriter, _tb_available = optional_import("torch.utils.tensorboard", name="SummaryWriter") SummaryWriterX, _ = optional_import("tensorboardX", name="SummaryWriter") DEFAULT_TAG = "Loss" @@ -46,10 +48,18 @@ class TensorBoardHandler: default to create a new TensorBoard writer. log_dir: if using default SummaryWriter, write logs to this directory, default is `./runs`. + Raises: + RuntimeError: When ``summary_writer`` is ``None`` and the ``tensorboard`` package is not installed. + """ def __init__(self, summary_writer: SummaryWriter | SummaryWriterX | None = None, log_dir: str = "./runs"): if summary_writer is None: + if not _tb_available: + raise RuntimeError( + "TensorBoardHandler requires tensorboard to be installed. " + "Please install it with: pip install tensorboard" + ) self._writer = SummaryWriter(log_dir=log_dir) self.internal_writer = True else: diff --git a/tests/handlers/test_handler_tb_stats.py b/tests/handlers/test_handler_tb_stats.py index b96bea13a1f..1b25b945f1a 100644 --- a/tests/handlers/test_handler_tb_stats.py +++ b/tests/handlers/test_handler_tb_stats.py @@ -14,12 +14,13 @@ import glob import tempfile import unittest -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch from ignite.engine import Engine, Events from parameterized import parameterized from monai.handlers import TensorBoardStatsHandler +from monai.handlers.tensorboard_handlers import TensorBoardHandler from monai.utils import optional_import SummaryWriter, has_tb = optional_import("torch.utils.tensorboard", name="SummaryWriter") @@ -162,5 +163,13 @@ def _update_metric(engine): ) # 2 = len([1, 3]) from event_filter +class TestTensorBoardHandlerMissingDependency(unittest.TestCase): + + def test_raises_when_tensorboard_unavailable(self): + with patch("monai.handlers.tensorboard_handlers._tb_available", False): + with self.assertRaises(RuntimeError): + TensorBoardHandler() + + if __name__ == "__main__": unittest.main()