from __future__ import annotations import torch import torch.nn as nn import torch.nn.functional as F def dice_loss_from_logits(logits: torch.Tensor, target: torch.Tensor, eps: float = 1.0) -> torch.Tensor: pred = torch.sigmoid(logits) dims = (1, 2, 3) inter = (pred * target).sum(dim=dims) denom = pred.sum(dim=dims) + target.sum(dim=dims) return (1.0 - (2.0 * inter + eps) / (denom + eps)).mean() def soft_erode(x: torch.Tensor) -> torch.Tensor: vertical = -F.max_pool2d(-x, kernel_size=(3, 1), stride=1, padding=(1, 0)) horizontal = -F.max_pool2d(-x, kernel_size=(1, 3), stride=1, padding=(0, 1)) return torch.minimum(vertical, horizontal) def soft_dilate(x: torch.Tensor) -> torch.Tensor: return F.max_pool2d(x, kernel_size=3, stride=1, padding=1) def soft_open(x: torch.Tensor) -> torch.Tensor: return soft_dilate(soft_erode(x)) def soft_skeletonize(x: torch.Tensor, iterations: int = 20) -> torch.Tensor: skeleton = F.relu(x - soft_open(x)) eroded = x for _ in range(max(int(iterations), 1)): eroded = soft_erode(eroded) opened = soft_open(eroded) delta = F.relu(eroded - opened) skeleton = skeleton + F.relu(delta - skeleton * delta) return skeleton.clamp(0, 1) def cldice_loss_from_logits(logits: torch.Tensor, target: torch.Tensor, iterations: int = 20, eps: float = 1.0) -> torch.Tensor: pred = torch.sigmoid(logits).clamp(0, 1) target = target.clamp(0, 1) pred_skel = soft_skeletonize(pred, iterations) target_skel = soft_skeletonize(target, iterations) dims = (1, 2, 3) topology_precision = ((pred_skel * target).sum(dim=dims) + eps) / (pred_skel.sum(dim=dims) + eps) topology_sensitivity = ((target_skel * pred).sum(dim=dims) + eps) / (target_skel.sum(dim=dims) + eps) cldice = (2.0 * topology_precision * topology_sensitivity) / (topology_precision + topology_sensitivity + 1e-6) return (1.0 - cldice).mean() class SegmentationLoss(nn.Module): def __init__( self, bce_weight: float = 1.0, dice_weight: float = 1.0, cldice_weight: float = 0.5, skeleton_iterations: int = 20, ): super().__init__() self.bce_weight = bce_weight self.dice_weight = dice_weight self.cldice_weight = cldice_weight self.skeleton_iterations = skeleton_iterations def forward(self, logits: torch.Tensor, target: torch.Tensor) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: bce = F.binary_cross_entropy_with_logits(logits, target) dice = dice_loss_from_logits(logits, target) cldice = cldice_loss_from_logits(logits, target, self.skeleton_iterations) loss = self.bce_weight * bce + self.dice_weight * dice + self.cldice_weight * cldice return loss, {"loss": loss.detach(), "bce_loss": bce.detach(), "dice_loss": dice.detach(), "cldice_loss": cldice.detach()}