Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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
1 change: 1 addition & 0 deletions clinicadl/losses/config/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,3 +2,4 @@

from .configs import *
from .enum import ImplementedLoss
from .monai import *
4 changes: 2 additions & 2 deletions clinicadl/losses/config/configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,9 +56,9 @@ def get_object(self) -> torch.nn.Module:
The PyTorch loss function.
"""
params = self.to_raw_dict()
if "weight" in params and params["weight"]:
if isinstance(params.get("weight"), list):
params["weight"] = torch.Tensor(params["weight"])
if "pos_weight" in params and params["pos_weight"]:
if isinstance(params.get("pos_weight"), list):
params["pos_weight"] = torch.Tensor(params["pos_weight"])

associated_class = self._get_class()
Expand Down
16 changes: 16 additions & 0 deletions clinicadl/losses/config/enum.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,14 @@ class ImplementedLoss(str, Enum):
SMOOTH_L1 = "SmoothL1Loss"
KLDIV = "KLDivLoss"

DICE = "DiceLoss"
DICE_CE = "DiceCELoss"
DICE_FOCAL = "DiceFocalLoss"
GENERALIZED_DICE = "GeneralizedDiceLoss"
GENERALIZED_DICE_FOCAL = "GeneralizedDiceFocalLoss"
FOCAL = "FocalLoss"
TVERSKY = "TverskyLoss"

@classmethod
def _missing_(cls, value):
raise ValueError(
Expand All @@ -36,3 +44,11 @@ class Order(int, Enum):

ONE = 1
TWO = 2


class GeneralizedDiceWeight(str, Enum):
"""Supported class weighting modes for generalized Dice losses."""

SQUARE = "square"
SIMPLE = "simple"
UNIFORM = "uniform"
166 changes: 166 additions & 0 deletions clinicadl/losses/config/monai.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,166 @@
"""Config classes for commonly used MONAI loss functions."""

from typing import Callable, Optional, Union

import monai.losses
import torch
from pydantic import NonNegativeFloat, field_validator, model_validator

from clinicadl.utils.factories import get_defaults_from

from .configs import LossConfig
from .enum import GeneralizedDiceWeight, Reduction

__all__ = [
"MonaiLossConfig",
"DiceLossConfig",
"DiceCELossConfig",
"DiceFocalLossConfig",
"GeneralizedDiceLossConfig",
"GeneralizedDiceFocalLossConfig",
"FocalLossConfig",
"TverskyLossConfig",
]

DICE_MONAI_DEFAULTS = get_defaults_from(monai.losses.DiceLoss)
DICE_CE_MONAI_DEFAULTS = get_defaults_from(monai.losses.DiceCELoss)
DICE_FOCAL_MONAI_DEFAULTS = get_defaults_from(monai.losses.DiceFocalLoss)
GENERALIZED_DICE_MONAI_DEFAULTS = get_defaults_from(monai.losses.GeneralizedDiceLoss)
GENERALIZED_DICE_FOCAL_MONAI_DEFAULTS = get_defaults_from(
monai.losses.GeneralizedDiceFocalLoss
)
FOCAL_MONAI_DEFAULTS = get_defaults_from(monai.losses.FocalLoss)
TVERSKY_MONAI_DEFAULTS = get_defaults_from(monai.losses.TverskyLoss)

SerializableWeight = Optional[Union[NonNegativeFloat, list[NonNegativeFloat]]]


class MonaiLossConfig(LossConfig):
"""Base config class for MONAI loss functions."""

@classmethod
def _get_class(cls) -> type[torch.nn.Module]:
"""Returns the MONAI loss function associated to this config class."""
return getattr(monai.losses, cls._get_name())


class _OverlapLossConfig(MonaiLossConfig):
"""Parameters shared by overlap-based MONAI losses."""

include_background: bool = DICE_MONAI_DEFAULTS["include_background"]
to_onehot_y: bool = DICE_MONAI_DEFAULTS["to_onehot_y"]
sigmoid: bool = DICE_MONAI_DEFAULTS["sigmoid"]
softmax: bool = DICE_MONAI_DEFAULTS["softmax"]
other_act: Optional[Callable[[torch.Tensor], torch.Tensor]] = DICE_MONAI_DEFAULTS[
"other_act"
]
reduction: Reduction = DICE_MONAI_DEFAULTS["reduction"]
smooth_nr: NonNegativeFloat = DICE_MONAI_DEFAULTS["smooth_nr"]
smooth_dr: NonNegativeFloat = DICE_MONAI_DEFAULTS["smooth_dr"]
batch: bool = DICE_MONAI_DEFAULTS["batch"]

@model_validator(mode="after")
def validate_activation(self):
"""Only one activation may be enabled at a time."""
if sum((self.sigmoid, self.softmax, self.other_act is not None)) > 1:
raise ValueError(
"Only one of 'sigmoid', 'softmax' and 'other_act' may be set."
)
return self


class _DiceLossConfig(_OverlapLossConfig):
"""Parameters shared by Dice-based MONAI losses."""

squared_pred: bool = DICE_MONAI_DEFAULTS["squared_pred"]
jaccard: bool = DICE_MONAI_DEFAULTS["jaccard"]


class DiceLossConfig(_DiceLossConfig):
"""Config class for :py:class:`monai.losses.DiceLoss`."""

weight: SerializableWeight = DICE_MONAI_DEFAULTS["weight"]
soft_label: bool = DICE_MONAI_DEFAULTS["soft_label"]


class DiceCELossConfig(_DiceLossConfig):
"""Config class for :py:class:`monai.losses.DiceCELoss`."""

weight: Optional[list[NonNegativeFloat]] = DICE_CE_MONAI_DEFAULTS["weight"]
lambda_dice: NonNegativeFloat = DICE_CE_MONAI_DEFAULTS["lambda_dice"]
lambda_ce: NonNegativeFloat = DICE_CE_MONAI_DEFAULTS["lambda_ce"]
label_smoothing: NonNegativeFloat = DICE_CE_MONAI_DEFAULTS["label_smoothing"]

@field_validator("label_smoothing")
@classmethod
def validate_label_smoothing(cls, value):
if value > 1:
raise ValueError("'label_smoothing' must be between 0 and 1.")
return value


class DiceFocalLossConfig(_DiceLossConfig):
"""Config class for :py:class:`monai.losses.DiceFocalLoss`."""

gamma: NonNegativeFloat = DICE_FOCAL_MONAI_DEFAULTS["gamma"]
weight: SerializableWeight = DICE_FOCAL_MONAI_DEFAULTS["weight"]
lambda_dice: NonNegativeFloat = DICE_FOCAL_MONAI_DEFAULTS["lambda_dice"]
lambda_focal: NonNegativeFloat = DICE_FOCAL_MONAI_DEFAULTS["lambda_focal"]
alpha: Optional[NonNegativeFloat] = DICE_FOCAL_MONAI_DEFAULTS["alpha"]

@field_validator("alpha")
@classmethod
def validate_alpha(cls, value):
if value is not None and value > 1:
raise ValueError("'alpha' must be between 0 and 1.")
return value


class _GeneralizedDiceLossConfig(_OverlapLossConfig):
"""Parameters shared by generalized Dice MONAI losses."""

w_type: GeneralizedDiceWeight = GENERALIZED_DICE_MONAI_DEFAULTS["w_type"]


class GeneralizedDiceLossConfig(_GeneralizedDiceLossConfig):
"""Config class for :py:class:`monai.losses.GeneralizedDiceLoss`."""

soft_label: bool = GENERALIZED_DICE_MONAI_DEFAULTS["soft_label"]


class GeneralizedDiceFocalLossConfig(_GeneralizedDiceLossConfig):
"""Config class for :py:class:`monai.losses.GeneralizedDiceFocalLoss`."""

gamma: NonNegativeFloat = GENERALIZED_DICE_FOCAL_MONAI_DEFAULTS["gamma"]
weight: SerializableWeight = GENERALIZED_DICE_FOCAL_MONAI_DEFAULTS["weight"]
lambda_gdl: NonNegativeFloat = GENERALIZED_DICE_FOCAL_MONAI_DEFAULTS["lambda_gdl"]
lambda_focal: NonNegativeFloat = GENERALIZED_DICE_FOCAL_MONAI_DEFAULTS[
"lambda_focal"
]


class FocalLossConfig(MonaiLossConfig):
"""Config class for :py:class:`monai.losses.FocalLoss`."""

include_background: bool = FOCAL_MONAI_DEFAULTS["include_background"]
to_onehot_y: bool = FOCAL_MONAI_DEFAULTS["to_onehot_y"]
gamma: NonNegativeFloat = FOCAL_MONAI_DEFAULTS["gamma"]
alpha: Optional[NonNegativeFloat] = FOCAL_MONAI_DEFAULTS["alpha"]
weight: SerializableWeight = FOCAL_MONAI_DEFAULTS["weight"]
reduction: Reduction = FOCAL_MONAI_DEFAULTS["reduction"]
use_softmax: bool = FOCAL_MONAI_DEFAULTS["use_softmax"]

@field_validator("alpha")
@classmethod
def validate_alpha(cls, value):
if value is not None and value > 1:
raise ValueError("'alpha' must be between 0 and 1.")
return value


class TverskyLossConfig(_OverlapLossConfig):
"""Config class for :py:class:`monai.losses.TverskyLoss`."""

alpha: NonNegativeFloat = TVERSKY_MONAI_DEFAULTS["alpha"]
beta: NonNegativeFloat = TVERSKY_MONAI_DEFAULTS["beta"]
soft_label: bool = TVERSKY_MONAI_DEFAULTS["soft_label"]
19 changes: 18 additions & 1 deletion docs/api/losses.rst
Original file line number Diff line number Diff line change
Expand Up @@ -40,4 +40,21 @@ Regression / Reconstruction
L1LossConfig
SmoothL1LossConfig
HuberLossConfig
KLDivLossConfig
KLDivLossConfig


MONAI Segmentation
^^^^^^^^^^^^^^^^^^

.. autosummary::
:toctree: ../generated/
:nosignatures:
:template: autosummary/config_object.rst

DiceLossConfig
DiceCELossConfig
DiceFocalLossConfig
GeneralizedDiceLossConfig
GeneralizedDiceFocalLossConfig
FocalLossConfig
TverskyLossConfig
66 changes: 66 additions & 0 deletions tests/unittests/losses/test_config.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,18 @@
import pytest
import torch.nn as nn
from monai import losses as monai_losses
from pydantic import ValidationError

from clinicadl.losses.config import (
BCELossConfig,
BCEWithLogitsLossConfig,
CrossEntropyLossConfig,
DiceCELossConfig,
DiceFocalLossConfig,
DiceLossConfig,
FocalLossConfig,
GeneralizedDiceFocalLossConfig,
GeneralizedDiceLossConfig,
HuberLossConfig,
ImplementedLoss,
KLDivLossConfig,
Expand All @@ -14,6 +21,7 @@
MultiMarginLossConfig,
NLLLossConfig,
SmoothL1LossConfig,
TverskyLossConfig,
)

BAD_INPUTS = [
Expand Down Expand Up @@ -172,3 +180,61 @@ def test_name():
config = globals()[f"{name.value}Config"]
c = config()
assert c.name_ == name.value


@pytest.mark.parametrize(
"config,loss",
[
(DiceLossConfig, monai_losses.DiceLoss),
(DiceCELossConfig, monai_losses.DiceCELoss),
(DiceFocalLossConfig, monai_losses.DiceFocalLoss),
(GeneralizedDiceLossConfig, monai_losses.GeneralizedDiceLoss),
(
GeneralizedDiceFocalLossConfig,
monai_losses.GeneralizedDiceFocalLoss,
),
(FocalLossConfig, monai_losses.FocalLoss),
(TverskyLossConfig, monai_losses.TverskyLoss),
],
)
def test_get_monai_object(config, loss):
assert isinstance(config().get_object(), loss)


@pytest.mark.parametrize(
"config",
[
DiceLossConfig,
DiceCELossConfig,
DiceFocalLossConfig,
GeneralizedDiceLossConfig,
GeneralizedDiceFocalLossConfig,
TverskyLossConfig,
],
)
def test_monai_loss_rejects_multiple_activations(config):
with pytest.raises(ValidationError, match="Only one"):
config(sigmoid=True, softmax=True)


@pytest.mark.parametrize(
"config,kwargs",
[
(DiceCELossConfig, {"label_smoothing": 1.1}),
(DiceFocalLossConfig, {"alpha": 1.1}),
(FocalLossConfig, {"alpha": 1.1}),
(DiceFocalLossConfig, {"gamma": -1}),
(GeneralizedDiceFocalLossConfig, {"lambda_gdl": -1}),
(GeneralizedDiceLossConfig, {"w_type": "invalid"}),
(TverskyLossConfig, {"smooth_dr": -1}),
],
)
def test_bad_monai_inputs(config, kwargs):
with pytest.raises(ValidationError):
config(**kwargs)


def test_monai_weight_list_is_converted_to_tensor():
loss = FocalLossConfig(weight=[1, 2]).get_object()
assert loss.class_weight is not None
assert loss.class_weight.tolist() == [1, 2]
19 changes: 19 additions & 0 deletions tests/unittests/losses/test_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,13 @@
MultiMarginLossConfig,
NLLLossConfig,
SmoothL1LossConfig,
DiceLossConfig,
DiceCELossConfig,
DiceFocalLossConfig,
GeneralizedDiceLossConfig,
GeneralizedDiceFocalLossConfig,
FocalLossConfig,
TverskyLossConfig,
],
)
def test_get_loss_function_from_dict(config, tmp_path):
Expand All @@ -30,3 +37,15 @@ def test_get_loss_function_from_dict(config, tmp_path):
if config is NLLLossConfig:
c = NLLLossConfig(weight=[1, 2])
assert get_loss_function_from_dict(c.to_dict()).weight == [1, 2]


@pytest.mark.parametrize(
"config",
[
DiceLossConfig(sigmoid=True, weight=[1, 2]),
GeneralizedDiceLossConfig(w_type="simple", soft_label=True),
FocalLossConfig(gamma=3, alpha=0.25, use_softmax=True),
],
)
def test_monai_loss_non_default_round_trip(config):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could you add a docstring to explain why this test is necessary?

assert get_loss_function_from_dict(config.to_dict()) == config
Loading