""" Lovász-Softmax loss — NaN-safe version for small-batch training. Fixes over v1: - eps guard on union denominator (prevents /0 on rare classes) - support for multiple ignore indices (background + void) """ import torch import torch.nn as nn import torch.nn.functional as F def lovasz_grad(gt_sorted): """Gradient of Lovász extension w.r.t. sorted errors.""" p = len(gt_sorted) gts = gt_sorted.sum() intersection = gts - gt_sorted.float().cumsum(0) union = gts + (1.0 - gt_sorted).float().cumsum(0) jaccard = 1.0 - intersection / (union + 1e-6) # eps guard if p > 1: jaccard[1:p] = jaccard[1:p] - jaccard[0:-1] return jaccard def lovasz_softmax_flat(probas, labels, classes="present"): """ Multi-class Lovász-Softmax on flattened tensors. probas: (N, C) softmax probabilities labels: (N,) ground truth indices """ if probas.numel() == 0: return probas * 0.0 C = probas.shape[1] if classes == "all": class_list = list(range(C)) elif classes == "present": class_list = torch.unique(labels).tolist() else: class_list = list(classes) losses = [] for c in class_list: fg = (labels == c).float() if fg.sum() == 0: continue class_pred = probas[:, c] errors = (fg - class_pred).abs() errors_sorted, perm = torch.sort(errors, dim=0, descending=True) fg_sorted = fg[perm.data] grad = lovasz_grad(fg_sorted) loss_c = torch.dot(errors_sorted, grad) if torch.isfinite(loss_c): losses.append(loss_c) if not losses: return probas.sum() * 0.0 return sum(losses) / len(losses) class LovaszSoftmaxLoss(nn.Module): """ Lovász-Softmax for 2D segmentation. logits: (B, C, H, W) targets: (B, H, W) long ignore_indices: list of class indices to exclude (e.g. [0, 19]) """ def __init__(self, ignore_indices=None, classes="present"): super().__init__() self.ignore_indices = ignore_indices or [] self.classes = classes def forward(self, logits, targets): probas = F.softmax(logits, dim=1) B, C, H, W = probas.shape probas_flat = probas.permute(0, 2, 3, 1).contiguous().view(-1, C) targets_flat = targets.contiguous().view(-1) # Mask out all ignore indices if self.ignore_indices: valid = torch.ones(targets_flat.shape[0], dtype=torch.bool, device=targets_flat.device) for idx in self.ignore_indices: valid &= (targets_flat != idx) probas_flat = probas_flat[valid] targets_flat = targets_flat[valid] if probas_flat.numel() == 0: return logits.sum() * 0.0 loss = lovasz_softmax_flat(probas_flat, targets_flat, classes=self.classes) return torch.nan_to_num(loss, nan=0.0)