Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
134 changes: 13 additions & 121 deletions src/virtual_stain_flow/vsf_logging/MlflowLogger.py
Comment thread
wli51 marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -34,6 +38,7 @@ class MlflowLogger:
"""
def __init__(
self,
*,
name: str,
experiment_name: str,
tracking_uri: Optional[path_type] = None,
Expand Down Expand Up @@ -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

"""
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand Down
9 changes: 9 additions & 0 deletions src/virtual_stain_flow/vsf_logging/auto_loggers/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
]
Original file line number Diff line number Diff line change
@@ -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
Original file line number Diff line number Diff line change
@@ -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}")
Loading
Loading