Spaces:
Running
Running
| 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()} | |