from monai.networks.nets import UNet from monai.losses import DiceCELoss from monai.metrics import DiceMetric from monai.transforms import Activations, AsDiscrete model = UNet( spatial_dims=2, in_channels=1, out_channels=1, channels=(16, 32, 64, 128, 256), strides=(2, 2, 2, 2), num_res_units=2, ).to(device) loss_fn = DiceCELoss(sigmoid=True) # Dice handles the imbalance; CE smooths the gradient metric = DiceMetric(include_background=True, reduction="mean") post_pred = Compose([Activations(sigmoid=True), AsDiscrete(threshold=0.5)])