diff --git a/src/virtual_stain_flow/vsf_logging/MlflowLogger.py b/src/virtual_stain_flow/vsf_logging/MlflowLogger.py index a9164da..8933283 100644 --- a/src/virtual_stain_flow/vsf_logging/MlflowLogger.py +++ b/src/virtual_stain_flow/vsf_logging/MlflowLogger.py @@ -10,8 +10,12 @@ import mlflow from torch import nn -from ..models.model import BaseModel from ..trainers.trainer_protocol import TrainerProtocol +from .auto_loggers import ( + AutoLossGroupConfigLogger, + AutoModelConfigLogger, + AutoOptimizerConfigLogger, +) from .callbacks.LoggerCallback import ( AbstractLoggerCallback, log_type @@ -34,6 +38,7 @@ class MlflowLogger: """ def __init__( self, + *, name: str, experiment_name: str, tracking_uri: Optional[path_type] = None, @@ -137,6 +142,10 @@ def __init__( self._save_model_every_n_epochs = save_model_every_n_epochs self._save_best_model = save_best_model + self._model_config_logger = AutoModelConfigLogger(self) + self._optimizer_config_logger = AutoOptimizerConfigLogger(self) + self._loss_group_config_logger = AutoLossGroupConfigLogger(self) + return None """ @@ -196,39 +205,9 @@ def on_train_start(self): self._run_id = mlflow.active_run().info.run_id - # log model config if available - if hasattr(self.trainer, '_models') and isinstance(self.trainer._models, List): - models = self.trainer._models - elif hasattr(self.trainer, 'model'): - models = [self.trainer.model] - else: - models = [] - - for idx, model in enumerate(models): - - if isinstance(model, BaseModel) and hasattr(model, 'to_config'): - try: - config = model.to_config() # type: ignore - except Exception as e: - print(f"Could not get model config for logging: {e}") - config = None - if config: - class_path = config.get("class_path") - if class_path: - mlflow.set_tag( - f"model.{idx}.class_path", - str(class_path) - ) - try: - self.log_config( - tag=model.__class__.__name__, - config=config, # type: ignore - stage=None - ) - except Exception as e: - print(f"Fail to log model config as artifact: {e}") - - self._log_loss_groups_config_and_tags() + self._model_config_logger.log_model_configs(self.trainer) + self._optimizer_config_logger.log_optimizer_configs(self.trainer) + self._loss_group_config_logger.log_loss_group_configs(self.trainer) for callback in self.callbacks: @@ -510,93 +489,6 @@ def _save_model_weights( artifact_path=artifact_path ) - def _get_loss_groups(self) -> Dict[str, Any]: - """ - Discover loss groups attached to the bound trainer. - """ - - if self.trainer is None: - return {} - - loss_groups: Dict[str, Any] = {} - - explicit_groups = getattr(self.trainer, 'loss_groups', None) - if isinstance(explicit_groups, dict): - for group_name, group in explicit_groups.items(): - if hasattr(group, 'get_config'): - loss_groups[str(group_name)] = group - - fallback_attrs = { - 'main': '_loss_group', - 'generator': '_generator_loss_group', - 'discriminator': '_discriminator_loss_group' - } - for group_name, attr in fallback_attrs.items(): - if group_name in loss_groups: - continue - group = getattr(self.trainer, attr, None) - if group is not None and hasattr(group, 'get_config'): - loss_groups[group_name] = group - - return loss_groups - - def _log_loss_groups_config_and_tags(self) -> None: - """ - Log loss item names and weights as flattened string mlflow tags and - full loss group configuration (loss name, loss weight, whether loss - is active during validation etc.) as mlflow config artifacts. - """ - - loss_groups = self._get_loss_groups() - if not loss_groups: - return None - - for group_name, group in loss_groups.items(): - try: - group_config = group.get_config() - except Exception as e: - print( - f"Could not get loss group config for logging " - f"({group_name}): {e}" - ) - continue - - if not isinstance(group_config, list): - continue - - for idx, item_cfg in enumerate(group_config): - if not isinstance(item_cfg, dict): - continue - - if 'key' in item_cfg and item_cfg['key'] is not None: - mlflow.set_tag( - f"loss.{group_name}.{idx}.name", - str(item_cfg['key']) - ) - - if 'weight' in item_cfg and item_cfg['weight'] is not None: - mlflow.set_tag( - f"loss.{group_name}.{idx}.weight", - str(item_cfg['weight']) - ) - - try: - self.log_config( - tag=f"loss_group_{group_name}", - config={ - 'group_name': group_name, - 'items': group_config - }, - stage=None - ) - except Exception as e: - print( - f"Fail to log loss group config as artifact " - f"({group_name}): {e}" - ) - - return None - def log_config( self, tag: str, diff --git a/src/virtual_stain_flow/vsf_logging/auto_loggers/__init__.py b/src/virtual_stain_flow/vsf_logging/auto_loggers/__init__.py new file mode 100644 index 0000000..2860937 --- /dev/null +++ b/src/virtual_stain_flow/vsf_logging/auto_loggers/__init__.py @@ -0,0 +1,9 @@ +from .loss_group_config_logger import AutoLossGroupConfigLogger +from .model_config_logger import AutoModelConfigLogger +from .optimizer_config_logger import AutoOptimizerConfigLogger + +__all__ = [ + "AutoModelConfigLogger", + "AutoOptimizerConfigLogger", + "AutoLossGroupConfigLogger", +] diff --git a/src/virtual_stain_flow/vsf_logging/auto_loggers/loss_group_config_logger.py b/src/virtual_stain_flow/vsf_logging/auto_loggers/loss_group_config_logger.py new file mode 100644 index 0000000..14d5fe6 --- /dev/null +++ b/src/virtual_stain_flow/vsf_logging/auto_loggers/loss_group_config_logger.py @@ -0,0 +1,97 @@ +from typing import Any, Dict, Optional + +import mlflow + +from ...trainers.trainer_protocol import TrainerProtocol + + +class AutoLossGroupConfigLogger: + """ + Auto-log loss group metadata to MLflow. + """ + + def __init__(self, logger: Any) -> None: + self._logger = logger + + def discover_loss_groups( + self, + trainer: Optional[TrainerProtocol], + ) -> Dict[str, Any]: + if trainer is None: + return {} + + loss_groups: Dict[str, Any] = {} + + explicit_groups = getattr(trainer, "loss_groups", None) + if isinstance(explicit_groups, dict): + for group_name, group in explicit_groups.items(): + if hasattr(group, "get_config"): + loss_groups[str(group_name)] = group + + fallback_attrs = { + "main": "_loss_group", + "generator": "_generator_loss_group", + "discriminator": "_discriminator_loss_group", + } + for group_name, attr in fallback_attrs.items(): + if group_name in loss_groups: + continue + group = getattr(trainer, attr, None) + if group is not None and hasattr(group, "get_config"): + loss_groups[group_name] = group + + return loss_groups + + def log_loss_group_configs( + self, + trainer: Optional[TrainerProtocol], + ) -> None: + loss_groups = self.discover_loss_groups(trainer) + if not loss_groups: + return None + + for group_name, group in loss_groups.items(): + try: + group_config = group.get_config() + except Exception as e: + print( + f"Could not get loss group config for logging " + f"({group_name}): {e}" + ) + continue + + if not isinstance(group_config, list): + continue + + for idx, item_cfg in enumerate(group_config): + if not isinstance(item_cfg, dict): + continue + + if "key" in item_cfg and item_cfg["key"] is not None: + mlflow.set_tag( + f"loss.{group_name}.{idx}.name", + str(item_cfg["key"]), + ) + + if "weight" in item_cfg and item_cfg["weight"] is not None: + mlflow.set_tag( + f"loss.{group_name}.{idx}.weight", + str(item_cfg["weight"]), + ) + + try: + self._logger.log_config( + tag=f"loss_group_{group_name}", + config={ + "group_name": group_name, + "items": group_config, + }, + stage=None, + ) + except Exception as e: + print( + f"Fail to log loss group config as artifact " + f"({group_name}): {e}" + ) + + return None diff --git a/src/virtual_stain_flow/vsf_logging/auto_loggers/model_config_logger.py b/src/virtual_stain_flow/vsf_logging/auto_loggers/model_config_logger.py new file mode 100644 index 0000000..bc8143a --- /dev/null +++ b/src/virtual_stain_flow/vsf_logging/auto_loggers/model_config_logger.py @@ -0,0 +1,61 @@ +from typing import Any, List, Optional + +import mlflow + +from ...models.model import BaseModel +from ...trainers.trainer_protocol import TrainerProtocol + + +class AutoModelConfigLogger: + """ + Auto-log model configuration metadata to MLflow. + """ + + def __init__(self, logger: Any) -> None: + self._logger = logger + + def _discover_models(self, trainer: Optional[TrainerProtocol]) -> List[Any]: + if trainer is None: + return [] + + explicit_models = getattr(trainer, "_models", None) + if isinstance(explicit_models, list): + return explicit_models + + model = getattr(trainer, "model", None) + if model is not None: + return [model] + + return [] + + def log_model_configs(self, trainer: Optional[TrainerProtocol]) -> None: + models = self._discover_models(trainer) + + for idx, model in enumerate(models): + if not isinstance(model, BaseModel) or not hasattr(model, "to_config"): + continue + + try: + config = model.to_config() + except Exception as e: + print(f"Could not get model config for logging: {e}") + config = None + + if not isinstance(config, dict): + continue + + class_path = config.get("class_path") + if class_path: + mlflow.set_tag( + f"model.{idx}.class_path", + str(class_path), + ) + + try: + self._logger.log_config( + tag=model.__class__.__name__, + config=config, + stage=None, + ) + except Exception as e: + print(f"Fail to log model config as artifact: {e}") diff --git a/src/virtual_stain_flow/vsf_logging/auto_loggers/optimizer_config_logger.py b/src/virtual_stain_flow/vsf_logging/auto_loggers/optimizer_config_logger.py new file mode 100644 index 0000000..fa3f344 --- /dev/null +++ b/src/virtual_stain_flow/vsf_logging/auto_loggers/optimizer_config_logger.py @@ -0,0 +1,73 @@ +from typing import Any, Dict, List, Optional + +import mlflow +from torch.optim import Optimizer + +from ...trainers.trainer_protocol import TrainerProtocol + + +class AutoOptimizerConfigLogger: + """ + Auto-log optimizer metadata to MLflow. + """ + + def __init__(self, logger: Any) -> None: + self._logger = logger + + def discover_optimizers( + self, + trainer: Optional[TrainerProtocol], + ) -> List[Optimizer]: + if trainer is None: + return [] + + optimizers: List[Optimizer] = [] + + explicit_optimizers = getattr(trainer, "optimizers", None) + if isinstance(explicit_optimizers, list): + optimizers.extend(explicit_optimizers) + + explicit_optimizer = getattr(trainer, "optimizer", None) + if explicit_optimizer is not None: + optimizers.append(explicit_optimizer) + + return optimizers + + def log_optimizer_configs( + self, + trainer: Optional[TrainerProtocol], + ) -> None: + optimizers = self.discover_optimizers(trainer) + + for idx, optimizer in enumerate(optimizers): + if not isinstance(optimizer, Optimizer): + continue + + try: + opt_config: Optional[Dict[str, Any]] = { + "class_path": ( + f"{optimizer.__class__.__module__}." + f"{optimizer.__class__.__name__}" + ), + "defaults": dict(optimizer.defaults), + } + except Exception as e: + print(f"Could not get optimizer config for logging: {e}") + opt_config = None + + if opt_config is None: + continue + + mlflow.set_tag( + f"optimizer.{idx}.class_path", + str(opt_config.get("class_path")), + ) + + try: + self._logger.log_config( + tag=f"optimizer_{idx}", + config=opt_config, + stage=None, + ) + except Exception as e: + print(f"Fail to log optimizer config as artifact: {e}") diff --git a/tests/vsf_logging/test_autolog_param.py b/tests/vsf_logging/test_autolog_param.py new file mode 100644 index 0000000..1661994 --- /dev/null +++ b/tests/vsf_logging/test_autolog_param.py @@ -0,0 +1,306 @@ +from types import SimpleNamespace + +import pytest +import torch + +from virtual_stain_flow.models.model import BaseModel +from virtual_stain_flow.vsf_logging.auto_loggers.model_config_logger import ( + AutoModelConfigLogger, +) +from virtual_stain_flow.vsf_logging.auto_loggers.optimizer_config_logger import ( + AutoOptimizerConfigLogger, +) +from virtual_stain_flow.vsf_logging.auto_loggers.loss_group_config_logger import ( + AutoLossGroupConfigLogger, +) + + +class _DummyLogger: + def __init__(self): + self.logged = [] + + def log_config(self, tag, config, stage=None): + self.logged.append( + { + "tag": tag, + "config": config, + "stage": stage, + } + ) + + +class _FailingLogger(_DummyLogger): + def log_config(self, tag, config, stage=None): + raise RuntimeError("forced log failure") + + +def _make_optimizer(): + model = torch.nn.Linear(4, 2) + return torch.optim.Adam(model.parameters(), lr=1e-3) + + +class _FakeLossGroup: + def __init__(self, config): + self._config = config + + def get_config(self): + return self._config + + +class _FakeModel(BaseModel): + def __init__(self, config): + super().__init__() + self._config = config + + def forward(self, x): + return x + + def to_config(self): + return self._config + + @classmethod + def from_config(cls, config): + return cls(config) + + +def test_discover_optimizers_supports_list_and_single(): + logger = _DummyLogger() + auto_logger = AutoOptimizerConfigLogger(logger) + + opt_a = _make_optimizer() + opt_b = _make_optimizer() + trainer = SimpleNamespace(optimizers=[opt_a], optimizer=opt_b) + + optimizers = auto_logger.discover_optimizers(trainer) + + assert optimizers == [opt_a, opt_b] + + +def test_discover_optimizers_returns_empty_for_none_trainer(): + logger = _DummyLogger() + auto_logger = AutoOptimizerConfigLogger(logger) + + assert auto_logger.discover_optimizers(None) == [] + + +def test_log_optimizer_configs_sets_class_path_tags_and_artifacts(monkeypatch): + logger = _DummyLogger() + auto_logger = AutoOptimizerConfigLogger(logger) + + captured_tags = {} + + def fake_set_tag(key, value): + captured_tags[key] = value + + monkeypatch.setattr( + "virtual_stain_flow.vsf_logging.auto_loggers.optimizer_config_logger.mlflow.set_tag", + fake_set_tag, + ) + + optimizer = _make_optimizer() + trainer = SimpleNamespace(optimizer=optimizer) + + auto_logger.log_optimizer_configs(trainer) + + assert "optimizer.0.class_path" in captured_tags + assert captured_tags["optimizer.0.class_path"].endswith("Adam") + + assert len(logger.logged) == 1 + assert logger.logged[0]["tag"] == "optimizer_0" + assert logger.logged[0]["config"]["class_path"].endswith("Adam") + assert logger.logged[0]["config"]["defaults"]["lr"] == pytest.approx(1e-3) + + +def test_log_optimizer_configs_skips_non_optimizer_entries(monkeypatch): + logger = _DummyLogger() + auto_logger = AutoOptimizerConfigLogger(logger) + + captured_tags = {} + + def fake_set_tag(key, value): + captured_tags[key] = value + + monkeypatch.setattr( + "virtual_stain_flow.vsf_logging.auto_loggers.optimizer_config_logger.mlflow.set_tag", + fake_set_tag, + ) + + trainer = SimpleNamespace(optimizers=["not-an-optimizer"]) + + auto_logger.log_optimizer_configs(trainer) + + assert captured_tags == {} + assert logger.logged == [] + + +def test_log_optimizer_configs_swallows_log_config_failures(monkeypatch): + logger = _FailingLogger() + auto_logger = AutoOptimizerConfigLogger(logger) + + def fake_set_tag(_key, _value): + return None + + monkeypatch.setattr( + "virtual_stain_flow.vsf_logging.auto_loggers.optimizer_config_logger.mlflow.set_tag", + fake_set_tag, + ) + + trainer = SimpleNamespace(optimizer=_make_optimizer()) + + # Should not raise despite logger.log_config raising. + auto_logger.log_optimizer_configs(trainer) + + +def test_discover_models_prefers_models_list_over_single_model(): + logger = _DummyLogger() + auto_logger = AutoModelConfigLogger(logger) + + model_a = _FakeModel({"class_path": "pkg.ModelA", "init": {}}) + model_b = _FakeModel({"class_path": "pkg.ModelB", "init": {}}) + trainer = SimpleNamespace(_models=[model_a], model=model_b) + + models = auto_logger._discover_models(trainer) + + assert models == [model_a] + + +def test_log_model_configs_sets_class_path_tag_and_artifact(monkeypatch): + logger = _DummyLogger() + auto_logger = AutoModelConfigLogger(logger) + + captured_tags = {} + + def fake_set_tag(key, value): + captured_tags[key] = value + + monkeypatch.setattr( + "virtual_stain_flow.vsf_logging.auto_loggers.model_config_logger.mlflow.set_tag", + fake_set_tag, + ) + + model = _FakeModel({"class_path": "virtual_stain_flow.models.unet.UNet", "init": {"depth": 4}}) + trainer = SimpleNamespace(model=model) + + auto_logger.log_model_configs(trainer) + + assert captured_tags["model.0.class_path"].endswith("UNet") + assert len(logger.logged) == 1 + assert logger.logged[0]["tag"] == "_FakeModel" + assert logger.logged[0]["config"]["init"]["depth"] == 4 + + +def test_log_model_configs_skips_non_dict_configs(monkeypatch): + logger = _DummyLogger() + auto_logger = AutoModelConfigLogger(logger) + + captured_tags = {} + + def fake_set_tag(key, value): + captured_tags[key] = value + + monkeypatch.setattr( + "virtual_stain_flow.vsf_logging.auto_loggers.model_config_logger.mlflow.set_tag", + fake_set_tag, + ) + + model = _FakeModel(["not", "a", "dict"]) + trainer = SimpleNamespace(model=model) + + auto_logger.log_model_configs(trainer) + + assert captured_tags == {} + assert logger.logged == [] + + +def test_log_model_configs_swallows_log_config_failures(monkeypatch): + logger = _FailingLogger() + auto_logger = AutoModelConfigLogger(logger) + + def fake_set_tag(_key, _value): + return None + + monkeypatch.setattr( + "virtual_stain_flow.vsf_logging.auto_loggers.model_config_logger.mlflow.set_tag", + fake_set_tag, + ) + + model = _FakeModel({"class_path": "pkg.Model", "init": {}}) + trainer = SimpleNamespace(model=model) + + # Should not raise despite logger.log_config raising. + auto_logger.log_model_configs(trainer) + + +def test_discover_loss_groups_supports_explicit_and_fallback_attrs(): + logger = _DummyLogger() + auto_logger = AutoLossGroupConfigLogger(logger) + + main_group = _FakeLossGroup([{"key": "MSELoss", "weight": 1.0}]) + gen_group = _FakeLossGroup([{"key": "L1Loss", "weight": 0.5}]) + trainer = SimpleNamespace( + loss_groups={"main": main_group}, + _generator_loss_group=gen_group, + ) + + loss_groups = auto_logger.discover_loss_groups(trainer) + + assert set(loss_groups.keys()) == {"main", "generator"} + assert loss_groups["main"] is main_group + assert loss_groups["generator"] is gen_group + + +def test_log_loss_group_configs_sets_tags_and_logs_config_artifact(monkeypatch): + logger = _DummyLogger() + auto_logger = AutoLossGroupConfigLogger(logger) + + captured_tags = {} + + def fake_set_tag(key, value): + captured_tags[key] = value + + monkeypatch.setattr( + "virtual_stain_flow.vsf_logging.auto_loggers.loss_group_config_logger.mlflow.set_tag", + fake_set_tag, + ) + + group_items = [ + {"key": "MSELoss", "weight": 1.0}, + {"key": "L1Loss", "weight": 0.25}, + {"key": None, "weight": None}, + "ignored-non-dict-item", + ] + trainer = SimpleNamespace(loss_groups={"main": _FakeLossGroup(group_items)}) + + auto_logger.log_loss_group_configs(trainer) + + assert captured_tags["loss.main.0.name"] == "MSELoss" + assert captured_tags["loss.main.0.weight"] == "1.0" + assert captured_tags["loss.main.1.name"] == "L1Loss" + assert captured_tags["loss.main.1.weight"] == "0.25" + + assert len(logger.logged) == 1 + assert logger.logged[0]["tag"] == "loss_group_main" + assert logger.logged[0]["config"]["group_name"] == "main" + assert logger.logged[0]["config"]["items"] == group_items + + +def test_log_loss_group_configs_skips_non_list_config(monkeypatch): + logger = _DummyLogger() + auto_logger = AutoLossGroupConfigLogger(logger) + + captured_tags = {} + + def fake_set_tag(key, value): + captured_tags[key] = value + + monkeypatch.setattr( + "virtual_stain_flow.vsf_logging.auto_loggers.loss_group_config_logger.mlflow.set_tag", + fake_set_tag, + ) + + trainer = SimpleNamespace(loss_groups={"main": _FakeLossGroup({"not": "a-list"})}) + + auto_logger.log_loss_group_configs(trainer) + + assert captured_tags == {} + assert logger.logged == []