TSRDA / main_method /code /losses.py
Dhruv1000's picture
Organize complete final models, all ablations, logs and checkpoints with visual guides (part 7)
71d64bb verified
Raw History Blame Contribute Delete
2.97 kB
"""
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)