Skip to content
Draft
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
53 changes: 51 additions & 2 deletions mfai/pytorch/lightning_modules/segmentation.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,23 @@
import warnings
from pathlib import Path
from typing import Any, Literal, Tuple
import inspect

import lightning.pytorch as pl
import pandas as pd
import torch
import torchmetrics as tm
from pytorch_lightning.utilities import rank_zero_only
from torch import Tensor
import torch.optim as torch_optims
import torch.optim.lr_scheduler as torch_schedulers

from mfai.pytorch.models.base import BaseModel

def filter_kwargs(func, kwargs):
sig = inspect.signature(func)
return {k: v for k, v in kwargs.items() if k in sig.parameters}

# define custom scalar in tensorboard, to have 2 lines on same graph
layout = {
"Check Overfit": {
Expand All @@ -25,6 +32,11 @@ def __init__(
model: BaseModel,
type_segmentation: Literal["binary", "multiclass", "multilabel", "regression"],
loss: torch.nn.modules.loss._Loss,
learning_rate: float,
optimizer_cls: str = "Adam",
optimizer_kwargs: dict | None = None,
lr_scheduler_cls: str | None = None,
lr_scheduler_kwargs: dict | None = None
) -> None:
"""A lightning module adapted for segmentation of weather images.

Expand Down Expand Up @@ -56,6 +68,12 @@ def __init__(
self.model.input_shape[1],
)

self.learning_rate = learning_rate
self.optimizer_cls = optimizer_cls
self.optimizer_kwargs = optimizer_kwargs or {}
self.lr_scheduler_cls = lr_scheduler_cls
self.lr_scheduler_kwargs = lr_scheduler_kwargs or {}

def get_hparams(self) -> dict:
"""Return the hparams we want to save in logger"""
hparams = dict(self.hparams)
Expand Down Expand Up @@ -156,8 +174,39 @@ def save_test_metrics_as_csv(self, df: pd.DataFrame) -> None:
########################################################################################
# OPTIMIZER #
########################################################################################
def configure_optimizers(self) -> torch.optim.Adam:
return torch.optim.Adam(self.parameters(), lr=0.001)
def configure_optimizers(self) -> torch.optim.Optimizer | dict:

# Define optimizer with all given hyperparameters
optim_class = getattr(torch_optims, self.optimizer_cls)
optimizer = optim_class(self.parameters(),
lr=self.learning_rate,
**self.optimizer_kwargs)

# If given, define the lr scheduler
if self.lr_scheduler_cls:

scheduler_class = getattr(torch_schedulers, self.lr_scheduler_cls)
# Filter kwargs to fit the constructor
# Arguments such as "monitor" are not passed in the constructor
# but are only needed for Lightning to know what to condition the
# step() function on for certain schedulers (eg ReduceLROnPlateau)
constructor_kwargs = filter_kwargs(scheduler_class, self.scheduler_kwargs)
scheduler = scheduler_class(optimizer, **constructor_kwargs)

lr_scheduler_config = {
"scheduler": scheduler,
"interval": "epoch",
"frequency": 1,
"monitor": self.scheduler_kwargs.get("monitor", "val_loss"),
"strict": True,
"name": self.scheduler_kwargs.get("name", None),
}

return {
"optimizer": optimizer,
"lr_scheduler": lr_scheduler_config,
}
return optimizer

########################################################################################
# SHARED STEPS #
Expand Down
Loading