segment-monograms / src /losses.py
Saranga7's picture
Deploy monogram segmentation demo
bcc432f verified
Raw
History Blame Contribute Delete
2.92 kB
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()}